Create an elementwise kernel from an expression. Performs dense reindexing: `temp_ids` are mapped to 0, 1, 2, ... in sorted order.
(&self, expr: PyExpr)
| 59 | /// Create an elementwise kernel from an expression. |
| 60 | /// Performs dense reindexing: `temp_ids` are mapped to 0, 1, 2, ... in sorted order. |
| 61 | fn elementwise(&self, expr: PyExpr) -> PyResult<PyKernelSpec> { |
| 62 | // Collect temp_ids referenced in the expression, sorted |
| 63 | let referenced_temp_ids = expr.inner.referenced_args(); |
| 64 | let mut seen_names = HashSet::new(); |
| 65 | |
| 66 | // Build dense mapping: temp_id -> dense_index |
| 67 | let mapping: HashMap<usize, usize> = referenced_temp_ids |
| 68 | .iter() |
| 69 | .enumerate() |
| 70 | .map(|(dense_idx, &temp_id)| (temp_id, dense_idx)) |
| 71 | .collect(); |
| 72 | |
| 73 | // Rewrite ArgRefs to dense indices |
| 74 | let rewritten_expr = expr.inner.rewrite_arg_refs(&mapping); |
| 75 | |
| 76 | // Build ArgSpecs with names from temp_args |
| 77 | let arg_specs: Vec<ArgSpec> = referenced_temp_ids |
| 78 | .iter() |
| 79 | .map(|temp_id| -> PyResult<ArgSpec> { |
| 80 | let name = expr.temp_args.get(temp_id).cloned().ok_or_else(|| { |
| 81 | PyErr::new::<pyo3::exceptions::PyValueError, _>(format!( |
| 82 | "missing temp arg metadata for arg index {temp_id}" |
| 83 | )) |
| 84 | })?; |
| 85 | |
| 86 | if !seen_names.insert(name.clone()) { |
| 87 | return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(format!( |
| 88 | "duplicate arg name: {name}" |
| 89 | ))); |
| 90 | } |
| 91 | |
| 92 | Ok(ArgSpec { |
| 93 | name, |
| 94 | dtype: DType::F64, |
| 95 | is_scalar: false, |
| 96 | }) |
| 97 | }) |
| 98 | .collect::<PyResult<Vec<_>>>()?; |
| 99 | |
| 100 | let inner = KernelSpec { |
| 101 | kind: KernelKind::Elementwise, |
| 102 | args: arg_specs, |
| 103 | output: crate::ir::kernel::TensorSpec { dtype: DType::F64 }, |
| 104 | expr: rewritten_expr, |
| 105 | }; |
| 106 | Ok(PyKernelSpec { inner }) |
| 107 | } |
| 108 | |
| 109 | /// Map a kernel spec over named arguments, returning a `MapSpec` for `rt.go()`. |
| 110 | #[pyo3(signature = (spec, **kwargs))] |