aprender_contracts_cli/commands/
equations.rs1use std::path::Path;
2use std::str::FromStr;
3
4use provable_contracts::latex::{latex_escape, math_to_latex};
5use provable_contracts::schema::{parse_contract, Contract, Equation};
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum OutputFormat {
10 Text,
12 Latex,
14 Ptx,
16 Asm,
18}
19
20impl FromStr for OutputFormat {
21 type Err = String;
22
23 fn from_str(s: &str) -> Result<Self, String> {
25 match s {
26 "text" => Ok(Self::Text),
27 "latex" => Ok(Self::Latex),
28 "ptx" => Ok(Self::Ptx),
29 "asm" => Ok(Self::Asm),
30 other => Err(format!(
31 "unknown format '{other}', expected 'text', 'latex', 'ptx', or 'asm'"
32 )),
33 }
34 }
35}
36
37pub fn run(path: &Path, format: OutputFormat) -> Result<(), Box<dyn std::error::Error>> {
39 let contract = parse_contract(path)?;
40 let name = path
41 .file_stem()
42 .and_then(|s| s.to_str())
43 .unwrap_or("unknown");
44
45 match format {
46 OutputFormat::Text => render_text(name, &contract.equations),
47 OutputFormat::Latex => render_latex(name, &contract.equations),
48 OutputFormat::Ptx => render_ptx(name, &contract),
49 OutputFormat::Asm => render_asm(name, &contract),
50 }
51
52 Ok(())
53}
54
55fn render_text(name: &str, equations: &std::collections::BTreeMap<String, Equation>) {
57 println!("Equations for {name}");
58 println!("{}", "=".repeat(40 + name.len()));
59 println!();
60
61 for (id, eq) in equations {
62 println!(" {id}");
63 println!(" formula: {}", eq.formula);
64 if let Some(ref dom) = eq.domain {
65 println!(" domain: {dom}");
66 }
67 if let Some(ref cod) = eq.codomain {
68 println!(" codomain: {cod}");
69 }
70 if !eq.invariants.is_empty() {
71 println!(" invariants:");
72 for inv in &eq.invariants {
73 println!(" - {inv}");
74 }
75 }
76 println!();
77 }
78}
79
80fn render_latex(name: &str, equations: &std::collections::BTreeMap<String, Equation>) {
82 let escaped_name = latex_escape(name);
83 println!("% Equations for {name}");
84 println!("\\section{{Equations: {escaped_name}}}");
85 println!();
86
87 for (id, eq) in equations {
88 let escaped_id = latex_escape(id);
89 let latex_formula = math_to_latex(&eq.formula);
90
91 println!("\\subsection{{{escaped_id}}}");
92 println!("\\label{{eq:{id}}}");
93 println!();
94 println!("\\begin{{equation}}");
95 println!(" {latex_formula}");
96 println!("\\end{{equation}}");
97 println!();
98
99 if let Some(ref dom) = eq.domain {
100 println!("\\textbf{{Domain:}} ${}$", math_to_latex(dom));
101 println!();
102 }
103
104 if let Some(ref cod) = eq.codomain {
105 println!("\\textbf{{Codomain:}} ${}$", math_to_latex(cod));
106 println!();
107 }
108
109 if !eq.invariants.is_empty() {
110 println!("\\textbf{{Invariants:}}");
111 println!("\\begin{{itemize}}");
112 for inv in &eq.invariants {
113 println!(" \\item ${}$", math_to_latex(inv));
114 }
115 println!("\\end{{itemize}}");
116 println!();
117 }
118 }
119}
120
121fn render_ptx(name: &str, contract: &Contract) {
123 let kernel = kernel_name(name);
124 render_header_comment("PTX kernel stub", name, contract);
125 println!(".version 8.5");
126 println!(".target sm_90");
127 println!(".address_size 64");
128 println!();
129 println!(".visible .entry {kernel}(");
130 println!(" .param .u64 input,");
131 println!(" .param .u64 output,");
132 println!(" .param .u32 n");
133 println!(")");
134 println!("{{");
135 println!(" .reg .u32 %tid, %n;");
136 println!(" .reg .u64 %in_ptr, %out_ptr;");
137 println!(" .reg .f32 %val, %acc;");
138 println!();
139 println!(" ld.param.u64 %in_ptr, [input];");
140 println!(" ld.param.u64 %out_ptr, [output];");
141 println!(" ld.param.u32 %n, [n];");
142 println!();
143 println!(" mov.u32 %tid, %ctaid.x;");
144 println!(" mad.lo.u32 %tid, %tid, %ntid.x, %tid.x;");
145 println!();
146 render_body_comments(contract);
147 println!(" ret;");
148 println!("}}");
149}
150
151fn render_asm(name: &str, contract: &Contract) {
153 let kernel = kernel_name(name);
154 let isa = detect_simd_isa(contract);
155 render_header_comment(&format!("x86-64 {} stub", isa.label), name, contract);
156 println!(".intel_syntax noprefix");
157 println!(".text");
158 println!(".globl {kernel}_{}", isa.suffix);
159 println!(".p2align 4");
160 println!();
161 println!("{kernel}_{}:", isa.suffix);
162 println!(" push rbp");
163 println!(" mov rbp, rsp");
164 println!(" // rdi = input ptr, rsi = output ptr, edx = n");
165 println!();
166 if isa.width > 0 {
167 let r = isa.reg_prefix;
168 println!(" // {} registers: {r} x {}", isa.label, isa.reg_count);
169 for i in 0..std::cmp::min(isa.reg_count, 4) {
170 println!(" vxorps {r}{i}, {r}{i}, {r}{i}");
171 }
172 println!();
173 }
174 render_body_comments(contract);
175 println!(" pop rbp");
176 println!(" ret");
177}
178
179fn render_header_comment(label: &str, name: &str, contract: &Contract) {
181 println!("//");
182 println!("// {label}: {name}");
183 println!("// {}", contract.metadata.description);
184 for (id, eq) in &contract.equations {
185 println!("// Equation {id}: {}", eq.formula);
186 }
187 println!("//");
188 println!();
189}
190
191fn render_body_comments(contract: &Contract) {
193 if let Some(ref ks) = contract.kernel_structure {
194 for (i, phase) in ks.phases.iter().enumerate() {
195 println!(" // Phase {}: {}", i + 1, phase.name);
196 println!(" // {}", phase.description);
197 if let Some(ref inv) = phase.invariant {
198 println!(" // Invariant: {inv}");
199 }
200 println!();
201 }
202 } else {
203 for (id, eq) in &contract.equations {
204 println!(" // Equation: {id}");
205 println!(" // {}", eq.formula);
206 println!();
207 }
208 }
209 if !contract.proof_obligations.is_empty() {
210 println!(" // Proof obligations:");
211 for ob in &contract.proof_obligations {
212 println!(" // [{}] {}", ob.obligation_type, ob.property);
213 }
214 println!();
215 }
216}
217
218struct SimdIsa {
220 label: &'static str,
222 suffix: &'static str,
224 reg_prefix: &'static str,
226 reg_count: u32,
228 width: u32,
230}
231
232fn detect_simd_isa(contract: &Contract) -> SimdIsa {
234 let has = |pat: &str| {
235 contract
236 .simd_dispatch
237 .values()
238 .any(|m| m.keys().any(|k| k.contains(pat)))
239 };
240 if has("avx512") || has("512") {
241 SimdIsa {
242 label: "AVX-512",
243 suffix: "avx512",
244 reg_prefix: "zmm",
245 reg_count: 32,
246 width: 512,
247 }
248 } else if has("avx2") {
249 SimdIsa {
250 label: "AVX2",
251 suffix: "avx2",
252 reg_prefix: "ymm",
253 reg_count: 16,
254 width: 256,
255 }
256 } else if !contract.simd_dispatch.is_empty() {
257 SimdIsa {
258 label: "SSE4.1",
259 suffix: "sse41",
260 reg_prefix: "xmm",
261 reg_count: 16,
262 width: 128,
263 }
264 } else {
265 SimdIsa {
266 label: "scalar",
267 suffix: "scalar",
268 reg_prefix: "xmm",
269 reg_count: 0,
270 width: 0,
271 }
272 }
273}
274
275fn kernel_name(contract_name: &str) -> String {
278 let mut s = contract_name.to_string();
279 if let Some(pos) = s.rfind("-v") {
280 if s[pos + 2..].chars().all(|c| c.is_ascii_digit()) {
281 s.truncate(pos);
282 }
283 }
284 if let Some(stripped) = s.strip_suffix("-kernel") {
285 s = stripped.to_string();
286 }
287 s.replace('-', "_")
288}
289
290#[cfg(test)]
291mod tests {
292 use super::*;
293
294 #[test]
295 fn test_output_format_from_str() {
296 assert_eq!(OutputFormat::from_str("text").unwrap(), OutputFormat::Text);
297 assert_eq!(
298 OutputFormat::from_str("latex").unwrap(),
299 OutputFormat::Latex
300 );
301 assert_eq!(OutputFormat::from_str("ptx").unwrap(), OutputFormat::Ptx);
302 assert_eq!(OutputFormat::from_str("asm").unwrap(), OutputFormat::Asm);
303 assert!(OutputFormat::from_str("json").is_err());
304 }
305
306 #[test]
307 fn test_from_str_other_format_returns_descriptive_error() {
308 let other = "csv";
310 let err = OutputFormat::from_str(other).unwrap_err();
311 assert!(
312 err.contains(other),
313 "error should include the unrecognized format"
314 );
315 assert!(err.contains("unknown format"));
316 }
317
318 #[test]
319 fn test_kernel_name() {
320 assert_eq!(kernel_name("softmax-kernel-v1"), "softmax");
321 assert_eq!(kernel_name("rmsnorm-kernel-v1"), "rmsnorm");
322 assert_eq!(kernel_name("flash-attention-v1"), "flash_attention");
323 assert_eq!(
324 kernel_name("model-config-algebra-v1"),
325 "model_config_algebra"
326 );
327 assert_eq!(kernel_name("silu-kernel-v2"), "silu");
328 }
329
330 #[test]
331 fn test_detect_simd_isa_avx2() {
332 use provable_contracts::schema::parse_contract_str;
333 let yaml = "metadata:\n version: '1.0'\n description: test\n\
334 equations:\n eq1:\n formula: 'y = x'\n\
335 simd_dispatch:\n k:\n scalar: s\n avx2: a\n";
336 let contract = parse_contract_str(yaml).unwrap();
337 let isa = detect_simd_isa(&contract);
338 assert_eq!(isa.suffix, "avx2");
339 assert_eq!(isa.width, 256);
340 }
341}