Skip to main content

nxpu_backend_rockchip/
lib.rs

1//! Rockchip RKNN NPU backend for NxPU.
2//!
3//! Delegates compilation to the ONNX backend with Rockchip RKNN-specific
4//! validation and RKNN Toolkit 2 conversion hints.
5
6use nxpu_analysis::analyze;
7use nxpu_backend_core::{
8    Backend, BackendError, BackendOptions, BackendOutput, Diagnostic, DiagnosticLevel, Precision,
9    PrecisionPolicy, validate_patterns,
10};
11use nxpu_backend_onnx::OnnxBackend;
12use nxpu_ir::Module;
13
14mod support;
15
16use support::RknnNpuSupport;
17
18/// Rockchip RKNN NPU backend with RKNN Toolkit 2 hints.
19#[derive(Debug)]
20pub struct RockchipBackend;
21
22impl Backend for RockchipBackend {
23    fn name(&self) -> &str {
24        "Rockchip RKNN NPU"
25    }
26
27    fn targets(&self) -> &[&str] {
28        &["rockchip", "rknn"]
29    }
30
31    fn preferred_precision(&self) -> Precision {
32        Precision::Int8
33    }
34
35    fn compile(
36        &self,
37        module: &Module,
38        opts: &BackendOptions,
39    ) -> Result<BackendOutput, BackendError> {
40        let mut op_names = Vec::new();
41        for (i, ep) in module.entry_points.iter().enumerate() {
42            match analyze::classify_entry_point(module, i) {
43                Ok(pattern) => op_names.extend(analyze::pattern_op_names(&pattern)),
44                Err(e) => {
45                    return Err(BackendError::Unsupported(format!(
46                        "entry point '{}': {e}",
47                        ep.name
48                    )));
49                }
50            }
51        }
52
53        let precision = resolve_precision(opts, self.preferred_precision());
54        let op_refs: Vec<&str> = op_names.iter().map(|s| s.as_str()).collect();
55        let mut diagnostics = validate_patterns(&RknnNpuSupport, &op_refs, precision);
56
57        let mut output = OnnxBackend.compile(module, opts)?;
58        diagnostics.extend(output.diagnostics);
59
60        diagnostics.push(Diagnostic {
61            level: DiagnosticLevel::Info,
62            message: "RKNN Toolkit 2 conversion:\n  \
63                      rknn = RKNN()\n  \
64                      rknn.load_onnx(\"output.onnx\")\n  \
65                      rknn.build(do_quantization=True, dataset=\"calibration.txt\")\n  \
66                      rknn.export_rknn(\"model.rknn\")"
67                .into(),
68        });
69        diagnostics.push(Diagnostic {
70            level: DiagnosticLevel::Info,
71            message: "Target platform: RK3588 (3 TOPS NPU)".into(),
72        });
73
74        output.diagnostics = diagnostics;
75        Ok(output)
76    }
77}
78
79fn resolve_precision(opts: &BackendOptions, preferred: Precision) -> Precision {
80    match opts.precision {
81        PrecisionPolicy::Explicit(p) => p,
82        PrecisionPolicy::Auto => preferred,
83        PrecisionPolicy::Keep => Precision::F32,
84    }
85}
86
87#[cfg(test)]
88mod tests {
89    use super::*;
90    use nxpu_backend_core::{BackendOptions, OutputContent};
91
92    #[test]
93    fn backend_metadata() {
94        let backend = RockchipBackend;
95        assert_eq!(backend.name(), "Rockchip RKNN NPU");
96        assert!(backend.targets().contains(&"rockchip"));
97        assert!(backend.targets().contains(&"rknn"));
98        assert_eq!(backend.preferred_precision(), Precision::Int8);
99    }
100
101    #[test]
102    fn compile_matmul_with_rknn_hints() {
103        let source = std::fs::read_to_string(concat!(
104            env!("CARGO_MANIFEST_DIR"),
105            "/../../examples/matmul.wgsl"
106        ))
107        .unwrap();
108        let module = nxpu_parser::parse(&source).unwrap();
109
110        let output = RockchipBackend
111            .compile(&module, &BackendOptions::default())
112            .unwrap();
113        assert_eq!(output.files.len(), 1);
114        assert_eq!(output.files[0].name, "output.onnx");
115        assert!(matches!(output.files[0].content, OutputContent::Binary(_)));
116
117        let messages: Vec<&str> = output
118            .diagnostics
119            .iter()
120            .map(|d| d.message.as_str())
121            .collect();
122        assert!(messages.iter().any(|m| m.contains("RKNN")));
123        assert!(messages.iter().any(|m| m.contains("RK3588")));
124    }
125
126    fn load_and_compile(example: &str, opts: &BackendOptions) -> BackendOutput {
127        let source = std::fs::read_to_string(format!(
128            "{}/../../examples/{example}.wgsl",
129            env!("CARGO_MANIFEST_DIR")
130        ))
131        .unwrap();
132        let module = nxpu_parser::parse(&source).unwrap();
133        RockchipBackend.compile(&module, opts).unwrap()
134    }
135
136    #[test]
137    fn compile_conv2d() {
138        let output = load_and_compile("conv2d", &BackendOptions::default());
139        assert_ne!(output.files.len(), 0);
140        for file in &output.files {
141            assert_ne!(file.content.len(), 0);
142        }
143    }
144
145    #[test]
146    fn compile_relu() {
147        let output = load_and_compile("relu", &BackendOptions::default());
148        assert_ne!(output.files.len(), 0);
149        for file in &output.files {
150            assert_ne!(file.content.len(), 0);
151        }
152    }
153
154    #[test]
155    fn compile_attention() {
156        let output = load_and_compile("attention", &BackendOptions::default());
157        assert_ne!(output.files.len(), 0);
158        for file in &output.files {
159            assert_ne!(file.content.len(), 0);
160        }
161    }
162
163    #[test]
164    fn resolve_precision_explicit_and_keep() {
165        let explicit_opts = BackendOptions {
166            precision: PrecisionPolicy::Explicit(Precision::F16),
167            ..BackendOptions::default()
168        };
169        assert_eq!(
170            resolve_precision(&explicit_opts, Precision::Int8),
171            Precision::F16
172        );
173
174        let keep_opts = BackendOptions {
175            precision: PrecisionPolicy::Keep,
176            ..BackendOptions::default()
177        };
178        assert_eq!(
179            resolve_precision(&keep_opts, Precision::Int8),
180            Precision::F32
181        );
182    }
183
184    #[test]
185    fn all_rknn_diagnostics() {
186        let output = load_and_compile("matmul", &BackendOptions::default());
187        let messages: Vec<&str> = output
188            .diagnostics
189            .iter()
190            .map(|d| d.message.as_str())
191            .collect();
192        assert!(messages.iter().any(|m| m.contains("RKNN")));
193        assert!(messages.iter().any(|m| m.contains("RK3588")));
194    }
195}