slat_flow joins ss_flow: max abs diff 9.3e-6, 700/700 params, against upstream running on CPU torch. All three SLAT checkpoints load clean and run (img2shape 512/1024 and imgshape2tex, the last taking 64 in-channels since shape is concatenated). The SLAT flows differ from ss_flow in two ways, both handled in the shared core: tokens are a SparseTensor's voxels so attention runs per batch item, and RoPE phases are NOT shipped - positions are the input's own coordinates, so they are derived at call time by rope_phases_from_coords. tests/oracle_slat.py keeps the CPU-torch oracle harness: it patches both the dense and sparse flash-attn kernels with SDPA equivalents. Use it rather than reasoning about correctness - it has already caught three bugs a perfect 700/700 key match did not. |
||
|---|---|---|
| .. | ||
| __init__.py | ||
| convert.py | ||
| slat_flow.py | ||
| ss_flow.py | ||