Skip to main content

aprender_contracts_cli/commands/
equations.rs

1use 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/// Supported output formats for equation rendering
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum OutputFormat {
10    /// Plain-text tabular output
11    Text,
12    /// LaTeX document fragment with equation environments
13    Latex,
14    /// NVIDIA PTX kernel stub with equation comments
15    Ptx,
16    /// x86-64 assembly stub with SIMD register setup
17    Asm,
18}
19
20impl FromStr for OutputFormat {
21    type Err = String;
22
23    /// Parse a format name string into an `OutputFormat` variant
24    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
37/// Execute the equations command, rendering contract equations in the given format
38pub 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
55/// Render equations as human-readable plain text
56fn 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
80/// Render equations as LaTeX sections with math environments
81fn 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
121/// Render a PTX kernel stub with equation header comments
122fn 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
151/// Render an x86-64 assembly stub with SIMD register initialization
152fn 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
179/// Emit a block comment with kernel description and equation summaries
180fn 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
191/// Emit inline comments for kernel phases or equation formulas
192fn 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
218/// SIMD instruction set descriptor for assembly generation
219struct SimdIsa {
220    /// Human-readable ISA name (e.g. "AVX-512")
221    label: &'static str,
222    /// Function name suffix (e.g. "avx512")
223    suffix: &'static str,
224    /// Register prefix for the ISA (e.g. "zmm")
225    reg_prefix: &'static str,
226    /// Number of available SIMD registers
227    reg_count: u32,
228    /// Bit width of SIMD registers (0 for scalar)
229    width: u32,
230}
231
232/// Detect the highest SIMD ISA level from the contract's `simd_dispatch` map
233fn 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
275/// Derive a kernel function name from the contract stem.
276/// Strips version suffix and `-kernel`, converts hyphens to underscores.
277fn 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        // Exercises the `other =>` catch-all arm in from_str
309        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}