Skip to main content

nxpu_opt/
layout.rs

1//! Memory layout optimization pass.
2//!
3//! Assigns optimal memory layouts to tensor global variables based on the
4//! target backend's preferred format, and tracks where layout conversions
5//! would be needed.
6
7use nxpu_ir::{MemoryLayout, Module, TypeInner};
8
9use crate::Pass;
10
11/// Returns the permutation needed to convert from one layout to another.
12///
13/// NHWC -> NCHW = [0, 3, 1, 2]  (move channels from last to second)
14/// NCHW -> NHWC = [0, 2, 3, 1]  (move channels from second to last)
15///
16/// Returns `None` if layouts are the same or conversion is not applicable
17/// (e.g. between `RowMajor` and a spatial layout).
18pub fn layout_permutation(from: MemoryLayout, to: MemoryLayout) -> Option<Vec<i64>> {
19    match (from, to) {
20        (MemoryLayout::Nhwc, MemoryLayout::Nchw) => Some(vec![0, 3, 1, 2]),
21        (MemoryLayout::Nchw, MemoryLayout::Nhwc) => Some(vec![0, 2, 3, 1]),
22        _ => None, // Same layout or not a spatial layout conversion
23    }
24}
25
26/// Reorder 4-D shape dimensions from one layout to another.
27///
28/// Given dimensions in `from` layout order, returns them in `to` layout order.
29/// For example, if `from` = NCHW and the dims are `[N, C, H, W]`, converting to
30/// NHWC gives `[N, H, W, C]`.
31///
32/// Non-4-D shapes are returned unchanged, as are shapes where no permutation
33/// applies (same layout, or non-spatial layout pair).
34pub fn reorder_dims(dims: &[i64], from: MemoryLayout, to: MemoryLayout) -> Vec<i64> {
35    if dims.len() != 4 {
36        return dims.to_vec();
37    }
38    match layout_permutation(from, to) {
39        Some(perm) => perm.iter().map(|&p| dims[p as usize]).collect(),
40        None => dims.to_vec(),
41    }
42}
43
44/// Information about a needed transpose at a graph boundary.
45#[derive(Debug, Clone)]
46pub struct TransposeRecord {
47    /// Name of the global variable that needs transposing.
48    pub var_name: String,
49    /// The permutation to apply (e.g. `[0, 3, 1, 2]`).
50    pub perm: Vec<i64>,
51    /// Whether this is an input transpose (before computation) or output (after).
52    pub is_input: bool,
53}
54
55/// Pass that detects layout mismatches at graph boundaries and records
56/// the transposes needed to convert between layouts.
57///
58/// This is a two-phase approach:
59/// 1. [`LayoutTransform`] assigns target layouts to uninitialized globals.
60/// 2. `TransposeInsertion` detects globals whose current layout differs from
61///    the target and updates them, so the backend can emit the appropriate
62///    transpose operations.
63///
64/// The backend can call [`layout_permutation`] with the original and target
65/// layouts to obtain the concrete permutation vector.
66#[derive(Debug)]
67pub struct TransposeInsertion {
68    /// The target layout that the backend expects.
69    pub target: MemoryLayout,
70}
71
72impl Pass for TransposeInsertion {
73    fn name(&self) -> &str {
74        "TransposeInsertion"
75    }
76
77    fn run(&self, module: &mut Module) -> bool {
78        let mut changed = false;
79
80        for (_handle, gv) in module.global_variables.iter_mut() {
81            let current_layout = match gv.layout {
82                Some(l) => l,
83                None => continue,
84            };
85
86            if current_layout == self.target {
87                continue;
88            }
89
90            // Only convert between spatial layouts (NHWC <-> NCHW).
91            if layout_permutation(current_layout, self.target).is_some() {
92                gv.layout = Some(self.target);
93                changed = true;
94            }
95        }
96
97        changed
98    }
99}
100
101/// Assigns a target memory layout to all tensor-typed global variables
102/// that do not already have a layout annotation.
103///
104/// This pass propagates the target layout uniformly. A more advanced version
105/// could perform per-tensor analysis to minimize conversion overhead.
106#[derive(Debug)]
107pub struct LayoutTransform {
108    /// The target layout to assign.
109    pub target: MemoryLayout,
110}
111
112impl Pass for LayoutTransform {
113    fn name(&self) -> &str {
114        "LayoutTransform"
115    }
116
117    fn run(&self, module: &mut Module) -> bool {
118        let mut changed = false;
119
120        for (_handle, gv) in module.global_variables.iter_mut() {
121            // Only apply to tensor-typed globals or array globals that
122            // represent tensor data (storage buffers).
123            if gv.layout.is_some() {
124                continue;
125            }
126
127            let is_tensor_like = match &module.types[gv.ty].inner {
128                TypeInner::Tensor { .. } => true,
129                TypeInner::Array { .. } => {
130                    // Storage arrays are treated as flat tensor buffers.
131                    matches!(gv.space, nxpu_ir::AddressSpace::Storage { .. })
132                }
133                _ => false,
134            };
135
136            if is_tensor_like {
137                gv.layout = Some(self.target);
138                changed = true;
139            }
140        }
141
142        changed
143    }
144}
145
146/// Counts the number of layout mismatches between global variables.
147///
148/// A mismatch occurs when two globals that feed into the same operation
149/// have different layouts. This is useful for diagnostics.
150pub fn count_layout_mismatches(module: &Module) -> usize {
151    let layouts: Vec<Option<MemoryLayout>> = module
152        .global_variables
153        .iter()
154        .map(|(_, gv)| gv.layout)
155        .collect();
156
157    let assigned: Vec<MemoryLayout> = layouts.into_iter().flatten().collect();
158    if assigned.is_empty() {
159        return 0;
160    }
161
162    let first = assigned[0];
163    assigned.iter().filter(|&&l| l != first).count()
164}
165
166/// Returns the preferred memory layout for a given backend target name.
167pub fn preferred_layout_for_target(target: &str) -> MemoryLayout {
168    match target {
169        "tflite" | "litert" | "arm-ethos" | "ethos-u" | "mediatek" | "rockchip" => {
170            MemoryLayout::Nhwc
171        }
172        "onnx" | "intel" | "amd" | "qualcomm" => MemoryLayout::Nchw,
173        "coreml" | "apple-ane" => MemoryLayout::Nhwc,
174        _ => MemoryLayout::RowMajor,
175    }
176}
177
178#[cfg(test)]
179mod tests {
180    use super::*;
181    use nxpu_ir::*;
182
183    fn make_tensor_module() -> Module {
184        let mut module = Module::default();
185
186        let f32_ty = module.types.insert(Type {
187            name: None,
188            inner: TypeInner::Scalar(Scalar::F32),
189        });
190        let array_f32 = module.types.insert(Type {
191            name: None,
192            inner: TypeInner::Array {
193                base: f32_ty,
194                size: ArraySize::Dynamic,
195                stride: 4,
196            },
197        });
198        let tensor_ty = module.types.insert(Type {
199            name: None,
200            inner: TypeInner::Tensor {
201                scalar: Scalar::F32,
202                shape: TensorShape {
203                    dims: vec![
204                        Dimension::Dynamic(Some("batch".into())),
205                        Dimension::Fixed(224),
206                        Dimension::Fixed(224),
207                        Dimension::Fixed(3),
208                    ],
209                },
210            },
211        });
212
213        // Storage array (treated as tensor buffer)
214        module.global_variables.append(GlobalVariable {
215            name: Some("weights".into()),
216            space: AddressSpace::Storage {
217                access: StorageAccess::LOAD,
218            },
219            binding: Some(ResourceBinding {
220                group: 0,
221                binding: 0,
222            }),
223            ty: array_f32,
224            init: None,
225            layout: None,
226        });
227
228        // Explicit tensor type
229        module.global_variables.append(GlobalVariable {
230            name: Some("input".into()),
231            space: AddressSpace::Storage {
232                access: StorageAccess::LOAD,
233            },
234            binding: Some(ResourceBinding {
235                group: 0,
236                binding: 1,
237            }),
238            ty: tensor_ty,
239            init: None,
240            layout: None,
241        });
242
243        // Uniform (should NOT get layout)
244        let u32_ty = module.types.insert(Type {
245            name: None,
246            inner: TypeInner::Scalar(Scalar::U32),
247        });
248        let params_ty = module.types.insert(Type {
249            name: Some("Params".into()),
250            inner: TypeInner::Struct {
251                members: vec![StructMember {
252                    name: Some("N".into()),
253                    ty: u32_ty,
254                    offset: 0,
255                }],
256                span: 4,
257            },
258        });
259        module.global_variables.append(GlobalVariable {
260            name: Some("params".into()),
261            space: AddressSpace::Uniform,
262            binding: Some(ResourceBinding {
263                group: 0,
264                binding: 2,
265            }),
266            ty: params_ty,
267            init: None,
268            layout: None,
269        });
270
271        module
272    }
273
274    #[test]
275    fn assigns_nhwc_layout() {
276        let mut module = make_tensor_module();
277        let pass = LayoutTransform {
278            target: MemoryLayout::Nhwc,
279        };
280        let changed = pass.run(&mut module);
281        assert!(changed);
282
283        let layouts: Vec<_> = module
284            .global_variables
285            .iter()
286            .map(|(_, gv)| gv.layout)
287            .collect();
288
289        // weights (storage array) and input (tensor) should have NHWC
290        assert_eq!(layouts[0], Some(MemoryLayout::Nhwc));
291        assert_eq!(layouts[1], Some(MemoryLayout::Nhwc));
292        // params (uniform struct) should have no layout
293        assert_eq!(layouts[2], None);
294    }
295
296    #[test]
297    fn idempotent() {
298        let mut module = make_tensor_module();
299        let pass = LayoutTransform {
300            target: MemoryLayout::Nchw,
301        };
302        pass.run(&mut module);
303        let changed = pass.run(&mut module);
304        assert!(!changed); // Already assigned
305    }
306
307    #[test]
308    fn does_not_overwrite_existing() {
309        let mut module = make_tensor_module();
310        // Pre-assign NCHW to first variable
311        for (_, gv) in module.global_variables.iter_mut() {
312            if gv.name.as_deref() == Some("weights") {
313                gv.layout = Some(MemoryLayout::Nchw);
314            }
315        }
316
317        let pass = LayoutTransform {
318            target: MemoryLayout::Nhwc,
319        };
320        let changed = pass.run(&mut module);
321        assert!(changed); // Changed the tensor
322
323        let layouts: Vec<_> = module
324            .global_variables
325            .iter()
326            .map(|(_, gv)| gv.layout)
327            .collect();
328
329        // weights keeps NCHW (not overwritten)
330        assert_eq!(layouts[0], Some(MemoryLayout::Nchw));
331        // input gets NHWC
332        assert_eq!(layouts[1], Some(MemoryLayout::Nhwc));
333    }
334
335    #[test]
336    fn preferred_layout_tflite() {
337        assert_eq!(preferred_layout_for_target("tflite"), MemoryLayout::Nhwc);
338        assert_eq!(preferred_layout_for_target("arm-ethos"), MemoryLayout::Nhwc);
339    }
340
341    #[test]
342    fn preferred_layout_onnx() {
343        assert_eq!(preferred_layout_for_target("onnx"), MemoryLayout::Nchw);
344        assert_eq!(preferred_layout_for_target("intel"), MemoryLayout::Nchw);
345    }
346
347    #[test]
348    fn count_mismatches() {
349        let mut module = make_tensor_module();
350
351        // No layouts assigned yet
352        assert_eq!(count_layout_mismatches(&module), 0);
353
354        // Assign same layout to all
355        let pass = LayoutTransform {
356            target: MemoryLayout::Nhwc,
357        };
358        pass.run(&mut module);
359        assert_eq!(count_layout_mismatches(&module), 0);
360
361        // Override one to create a mismatch
362        for (_, gv) in module.global_variables.iter_mut() {
363            if gv.name.as_deref() == Some("weights") {
364                gv.layout = Some(MemoryLayout::Nchw);
365            }
366        }
367        assert_eq!(count_layout_mismatches(&module), 1);
368    }
369
370    // ---- layout_permutation tests ----
371
372    #[test]
373    fn layout_permutation_nhwc_to_nchw() {
374        assert_eq!(
375            layout_permutation(MemoryLayout::Nhwc, MemoryLayout::Nchw),
376            Some(vec![0, 3, 1, 2])
377        );
378    }
379
380    #[test]
381    fn layout_permutation_nchw_to_nhwc() {
382        assert_eq!(
383            layout_permutation(MemoryLayout::Nchw, MemoryLayout::Nhwc),
384            Some(vec![0, 2, 3, 1])
385        );
386    }
387
388    #[test]
389    fn layout_permutation_same_layout() {
390        assert_eq!(
391            layout_permutation(MemoryLayout::Nhwc, MemoryLayout::Nhwc),
392            None
393        );
394        assert_eq!(
395            layout_permutation(MemoryLayout::Nchw, MemoryLayout::Nchw),
396            None
397        );
398    }
399
400    #[test]
401    fn layout_permutation_row_major() {
402        assert_eq!(
403            layout_permutation(MemoryLayout::RowMajor, MemoryLayout::Nhwc),
404            None
405        );
406        assert_eq!(
407            layout_permutation(MemoryLayout::Nhwc, MemoryLayout::RowMajor),
408            None
409        );
410    }
411
412    #[test]
413    fn layout_permutation_col_major() {
414        assert_eq!(
415            layout_permutation(MemoryLayout::ColMajor, MemoryLayout::Nchw),
416            None
417        );
418    }
419
420    // ---- reorder_dims tests ----
421
422    #[test]
423    fn reorder_dims_nchw_to_nhwc() {
424        // NCHW [1, 3, 224, 224] -> NHWC [1, 224, 224, 3]
425        let nchw = vec![1, 3, 224, 224];
426        let nhwc = reorder_dims(&nchw, MemoryLayout::Nchw, MemoryLayout::Nhwc);
427        assert_eq!(nhwc, vec![1, 224, 224, 3]);
428    }
429
430    #[test]
431    fn reorder_dims_nhwc_to_nchw() {
432        // NHWC [1, 224, 224, 3] -> NCHW [1, 3, 224, 224]
433        let nhwc = vec![1, 224, 224, 3];
434        let nchw = reorder_dims(&nhwc, MemoryLayout::Nhwc, MemoryLayout::Nchw);
435        assert_eq!(nchw, vec![1, 3, 224, 224]);
436    }
437
438    #[test]
439    fn reorder_dims_non_4d_passthrough() {
440        let dims = vec![1, 2, 3];
441        assert_eq!(
442            reorder_dims(&dims, MemoryLayout::Nhwc, MemoryLayout::Nchw),
443            vec![1, 2, 3]
444        );
445    }
446
447    #[test]
448    fn reorder_dims_same_layout_passthrough() {
449        let dims = vec![1, 3, 224, 224];
450        assert_eq!(
451            reorder_dims(&dims, MemoryLayout::Nchw, MemoryLayout::Nchw),
452            vec![1, 3, 224, 224]
453        );
454    }
455
456    #[test]
457    fn reorder_dims_5d_passthrough() {
458        let dims = vec![1, 2, 3, 4, 5];
459        assert_eq!(
460            reorder_dims(&dims, MemoryLayout::Nhwc, MemoryLayout::Nchw),
461            vec![1, 2, 3, 4, 5]
462        );
463    }
464
465    #[test]
466    fn reorder_dims_roundtrip() {
467        let original = vec![1, 224, 224, 3]; // NHWC
468        let nchw = reorder_dims(&original, MemoryLayout::Nhwc, MemoryLayout::Nchw);
469        let back = reorder_dims(&nchw, MemoryLayout::Nchw, MemoryLayout::Nhwc);
470        assert_eq!(back, original);
471    }
472
473    // ---- TransposeInsertion pass tests ----
474
475    #[test]
476    fn no_transpose_when_layout_matches() {
477        let mut module = make_tensor_module();
478        // Assign all to Nhwc
479        let pass = LayoutTransform {
480            target: MemoryLayout::Nhwc,
481        };
482        pass.run(&mut module);
483        // All are Nhwc now, so TransposeInsertion targeting Nhwc should be a no-op
484        let pass2 = TransposeInsertion {
485            target: MemoryLayout::Nhwc,
486        };
487        assert!(!pass2.run(&mut module));
488    }
489
490    #[test]
491    fn insert_transpose_nhwc_to_nchw() {
492        let mut module = make_tensor_module();
493        // Assign Nhwc layout
494        let pass = LayoutTransform {
495            target: MemoryLayout::Nhwc,
496        };
497        pass.run(&mut module);
498        // Now convert to Nchw
499        let pass2 = TransposeInsertion {
500            target: MemoryLayout::Nchw,
501        };
502        assert!(pass2.run(&mut module));
503        // All tensor globals should now have Nchw layout
504        for (_, gv) in module.global_variables.iter() {
505            if gv.layout.is_some() {
506                assert_eq!(gv.layout, Some(MemoryLayout::Nchw));
507            }
508        }
509    }
510
511    #[test]
512    fn insert_transpose_nchw_to_nhwc() {
513        let mut module = make_tensor_module();
514        let pass = LayoutTransform {
515            target: MemoryLayout::Nchw,
516        };
517        pass.run(&mut module);
518        let pass2 = TransposeInsertion {
519            target: MemoryLayout::Nhwc,
520        };
521        assert!(pass2.run(&mut module));
522        for (_, gv) in module.global_variables.iter() {
523            if gv.layout.is_some() {
524                assert_eq!(gv.layout, Some(MemoryLayout::Nhwc));
525            }
526        }
527    }
528
529    #[test]
530    fn transpose_insertion_idempotent() {
531        let mut module = make_tensor_module();
532        let pass = LayoutTransform {
533            target: MemoryLayout::Nhwc,
534        };
535        pass.run(&mut module);
536        let pass2 = TransposeInsertion {
537            target: MemoryLayout::Nchw,
538        };
539        pass2.run(&mut module);
540        // Running again should be a no-op
541        assert!(!pass2.run(&mut module));
542    }
543
544    #[test]
545    fn transpose_insertion_skips_non_spatial() {
546        let mut module = make_tensor_module();
547        // Assign RowMajor layout to all
548        for (_, gv) in module.global_variables.iter_mut() {
549            let is_tensor_like = match &module.types[gv.ty].inner {
550                TypeInner::Tensor { .. } => true,
551                TypeInner::Array { .. } => {
552                    matches!(gv.space, AddressSpace::Storage { .. })
553                }
554                _ => false,
555            };
556            if is_tensor_like {
557                gv.layout = Some(MemoryLayout::RowMajor);
558            }
559        }
560        // TransposeInsertion to Nchw should not change RowMajor layouts
561        // (no permutation exists for RowMajor -> Nchw)
562        let pass = TransposeInsertion {
563            target: MemoryLayout::Nchw,
564        };
565        assert!(!pass.run(&mut module));
566    }
567
568    #[test]
569    fn transpose_insertion_skips_unset_layout() {
570        let mut module = make_tensor_module();
571        // No layouts are assigned yet
572        let pass = TransposeInsertion {
573            target: MemoryLayout::Nchw,
574        };
575        assert!(!pass.run(&mut module));
576    }
577
578    #[test]
579    fn transpose_insertion_name() {
580        let pass = TransposeInsertion {
581            target: MemoryLayout::Nchw,
582        };
583        assert_eq!(pass.name(), "TransposeInsertion");
584    }
585
586    #[test]
587    fn transpose_record_clone() {
588        let record = TransposeRecord {
589            var_name: "input".into(),
590            perm: vec![0, 3, 1, 2],
591            is_input: true,
592        };
593        let cloned = record.clone();
594        assert_eq!(cloned.var_name, "input");
595        assert_eq!(cloned.perm, vec![0, 3, 1, 2]);
596        assert!(cloned.is_input);
597    }
598
599    #[test]
600    fn layout_conv2d_shape_nhwc_to_nchw() {
601        // Conv2D NHWC [1, 224, 224, 3] -> NCHW [1, 3, 224, 224]
602        let nhwc = [1i64, 224, 224, 3];
603        let result = reorder_dims(&nhwc, MemoryLayout::Nhwc, MemoryLayout::Nchw);
604        assert_eq!(result, vec![1, 3, 224, 224]);
605    }
606
607    #[test]
608    fn layout_pool_shape_nchw_to_nhwc() {
609        // Pool NCHW [1, 64, 56, 56] -> NHWC [1, 56, 56, 64]
610        let nchw = [1i64, 64, 56, 56];
611        let result = reorder_dims(&nchw, MemoryLayout::Nchw, MemoryLayout::Nhwc);
612        assert_eq!(result, vec![1, 56, 56, 64]);
613    }
614
615    #[test]
616    fn layout_batchnorm_shape_nhwc_to_nchw() {
617        // BatchNorm NHWC [1, 32, 32, 128] -> NCHW [1, 128, 32, 32]
618        let nhwc = [1i64, 32, 32, 128];
619        let result = reorder_dims(&nhwc, MemoryLayout::Nhwc, MemoryLayout::Nchw);
620        assert_eq!(result, vec![1, 128, 32, 32]);
621    }
622}