jit = JITTensorKernel(memorybudget=1000000)
~/jit_tensor_kernel.pyA just-in-time compiler for a mini tensor computation DSL. The pipeline:
matmul(A,relu(B))+softmax(C) into an AST with var, call, and binop nodesmemory_budget (in elements) to prevent materializing intermediates that exceed budgetcompile() returns a callable, __call__() does one-shot compile+run# Usage:
jit = JITTensorKernel(memory_budget=1_000_000)
fn = jit.compile("matmul(A,relu(B))+softmax(C)",
shapes={"A": (8,16), "B": (16,8), "C": (8,8)})
result = fn(A=A, B=B, C=C)
# Or one-shot:
result = jit_tensor("relu(A)+B", {"A": (4,), "B": (4,)}, A=np.array([...]), B=np.array([...]))
| Optimization | Mechanism |
|---|---|
| Loop Fusion | Adjacent element-wise ops (relu, add, mul, softmax) are fused into a single vectorized expression, materializing only the final result |
| Tiling | Large dimensions (>64) are split into tiles to improve cache behavior |
| Memory Budget | If an intermediate exceeds memory_budget elements, the fusion group is split |
| Fusion Boundaries | matmul and sum break fusion chains since they require full tensor inputs |
All 29 tests pass covering: - **Parser**: simple vars, nested calls, binary ops, parentheses, invalid input - **SSA**: lowering, shape inference for matmul, binop lowering - **Fusion**: element-wise chains fuse; matmul is boundary; memory budget controls splitting; softmax fuses - **Tiling**: large dims tile; small dims don't - **Emission**: valid Python/NumPy output - **End-to-end**: relu, matmul, softmax individually and combined (`matmul(A,relu(B))+softmax(C)`) - **Edge cases**: empty tensors (0-dim), 1D tensors, different shapes, subtraction, caching ``` TOTAL: 29 | PASS: 29 | FAIL: 0 ``` Memory budget verification: with a tiny budget (1000 elements), `A+B+C` on 100×100 tensors splits into 2 groups; with a large budget it fuses into 1 group. ### Edge Cases Tested - **Empty tensors** `(0, 5)` — result shape `(0, 5)` - **Scalar-like** 1D tensors — correct element-wise behavior - **Different shapes** — same expression with different shapes recompiles correctly - **Subtraction** — `relu(A)+softmax(B)-C` lowers to add + mul-by-(-1) - **Caching** — identical expressions reuse compiled functions ---
{"model": "claude-sonnet-4-20250514", "problem_class": "jit-tensor-kernel", "result": "passed", "tests": 29}~/jit_tensor_kernel.pyA just-in-time compiler for a mini tensor computation DSL. The pipeline:
matmul(A,relu(B))+softmax(C) into an AST with var, call, and binop nodesmemory_budget (in elements) to prevent materializing intermediates that exceed budgetcompile() returns a callable, __call__() does one-shot compile+run# Usage:
jit = JITTensorKernel(memory_budget=1_000_000)
fn = jit.compile("matmul(A,relu(B))+softmax(C)",
shapes={"A": (8,16), "B": (16,8), "C": (8,8)})
result = fn(A=A, B=B, C=C)
# Or one-shot:
result = jit_tensor("relu(A)+B", {"A": (4,), "B": (4,)}, A=np.array([...]), B=np.array([...]))
| Optimization | Mechanism |
|---|---|
| Loop Fusion | Adjacent element-wise ops (relu, add, mul, softmax) are fused into a single vectorized expression, materializing only the final result |
| Tiling | Large dimensions (>64) are split into tiles to improve cache behavior |
| Memory Budget | If an intermediate exceeds memory_budget elements, the fusion group is split |
| Fusion Boundaries | matmul and sum break fusion chains since they require full tensor inputs |
All 29 tests pass covering: - **Parser**: simple vars, nested calls, binary ops, parentheses, invalid input - **SSA**: lowering, shape inference for matmul, binop lowering - **Fusion**: element-wise chains fuse; matmul is boundary; memory budget controls splitting; softmax fuses - **Tiling**: large dims tile; small dims don't - **Emission**: valid Python/NumPy output - **End-to-end**: relu, matmul, softmax individually and combined (`matmul(A,relu(B))+softmax(C)`) - **Edge cases**: empty tensors (0-dim), 1D tensors, different shapes, subtraction, caching ``` TOTAL: 29 | PASS: 29 | FAIL: 0 ``` Memory budget verification: with a tiny budget (1000 elements), `A+B+C` on 100×100 tensors splits into 2 groups; with a large budget it fuses into 1 group. ### Edge Cases Tested - **Empty tensors** `(0, 5)` — result shape `(0, 5)` - **Scalar-like** 1D tensors — correct element-wise behavior - **Different shapes** — same expression with different shapes recompiles correctly - **Subtraction** — `relu(A)+softmax(B)-C` lowers to add + mul-by-(-1) - **Caching** — identical expressions reuse compiled functions ---
{"model": "claude-sonnet-4-20250514", "problem_class": "jit-tensor-kernel", "result": "passed", "tests": 29}