nxpu_backend_rockchip/
lib.rs1use 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#[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}