Skip to main content

lift_tensor/
dialect.rs

1use crate::ops::TensorOp;
2use lift_core::dialect::Dialect;
3
4#[derive(Debug)]
5pub struct TensorDialect;
6
7impl Dialect for TensorDialect {
8    fn name(&self) -> &str {
9        "tensor"
10    }
11
12    fn verify_op(
13        &self,
14        op_name: &str,
15        num_inputs: usize,
16        num_results: usize,
17    ) -> Result<(), String> {
18        let full_name = if op_name.starts_with("tensor.") {
19            op_name.to_string()
20        } else {
21            format!("tensor.{}", op_name)
22        };
23
24        match TensorOp::from_name(&full_name) {
25            Some(op) => {
26                let (min_inputs, max_inputs) = op.num_inputs();
27                if num_inputs < min_inputs || num_inputs > max_inputs {
28                    return Err(format!(
29                        "Operation {} expects {}-{} inputs, got {}",
30                        full_name, min_inputs, max_inputs, num_inputs
31                    ));
32                }
33                let _ = num_results;
34                Ok(())
35            }
36            None => Err(format!("Unknown tensor operation: {}", full_name)),
37        }
38    }
39}
40
41pub fn register_tensor_dialect(registry: &mut lift_core::dialect::DialectRegistry) {
42    registry.register(Box::new(TensorDialect));
43}