1use nxpu_ir::{MemoryLayout, Module, TypeInner};
8
9use crate::Pass;
10
11pub 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, }
24}
25
26pub 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#[derive(Debug, Clone)]
46pub struct TransposeRecord {
47 pub var_name: String,
49 pub perm: Vec<i64>,
51 pub is_input: bool,
53}
54
55#[derive(Debug)]
67pub struct TransposeInsertion {
68 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 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#[derive(Debug)]
107pub struct LayoutTransform {
108 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 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 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
146pub 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
166pub 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 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 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 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 assert_eq!(layouts[0], Some(MemoryLayout::Nhwc));
291 assert_eq!(layouts[1], Some(MemoryLayout::Nhwc));
292 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); }
306
307 #[test]
308 fn does_not_overwrite_existing() {
309 let mut module = make_tensor_module();
310 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); let layouts: Vec<_> = module
324 .global_variables
325 .iter()
326 .map(|(_, gv)| gv.layout)
327 .collect();
328
329 assert_eq!(layouts[0], Some(MemoryLayout::Nchw));
331 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 assert_eq!(count_layout_mismatches(&module), 0);
353
354 let pass = LayoutTransform {
356 target: MemoryLayout::Nhwc,
357 };
358 pass.run(&mut module);
359 assert_eq!(count_layout_mismatches(&module), 0);
360
361 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 #[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 #[test]
423 fn reorder_dims_nchw_to_nhwc() {
424 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 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]; 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 #[test]
476 fn no_transpose_when_layout_matches() {
477 let mut module = make_tensor_module();
478 let pass = LayoutTransform {
480 target: MemoryLayout::Nhwc,
481 };
482 pass.run(&mut module);
483 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 let pass = LayoutTransform {
495 target: MemoryLayout::Nhwc,
496 };
497 pass.run(&mut module);
498 let pass2 = TransposeInsertion {
500 target: MemoryLayout::Nchw,
501 };
502 assert!(pass2.run(&mut module));
503 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 assert!(!pass2.run(&mut module));
542 }
543
544 #[test]
545 fn transpose_insertion_skips_non_spatial() {
546 let mut module = make_tensor_module();
547 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 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 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 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 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 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}