Compare commits

...
Author SHA1 Message Date
geohot b162ab15da pad invalid work (glm) 2026-07-16 19:05:49 -07:00
4 changed files with 133 additions and 4 deletions
+122
View File
@@ -0,0 +1,122 @@
# CONTINUE.md: PAD with Invalid instead of 0
## Goal
Make the low-level `Ops.PAD` pad with `Invalid` instead of `0`, while keeping
the external `Tensor.pad` behavior unchanged.
## Changes made (all 3 files are modified, see `git diff`)
### 1. `tinygrad/schedule/indexing.py:92` — core change
`convert_pad_to_where_to_keep_behavior_local` now uses `UOp.const(x.dtype, Invalid)`
instead of `UOp.const(x.dtype, 0)` as the else value. This is what makes `Ops.PAD`
pad with Invalid.
### 2. `tinygrad/uop/symbolic.py:87-99` — Invalid propagation rules
Added two new rules to `pm_data_invalid` so that `where(invalid_gate, a, const_b)`
uses `b` (the const) in don't-care positions instead of poisoning to Invalid.
This is needed so that `_pad_constant`'s mask `where(pad(ones_bool), base, value)`
works — the mask is a `where(valid, True, Invalid)` gate, and the else `value`
is a const.
The rules are restricted to only match when the gate's valid value is a **const**
(`UPat.cvar("x")`), to distinguish pad masks (where valid=True, a const) from
gather masks (where valid=loaded_data, not a const). Without this restriction,
`test_tensor_index` breaks because gather masks also create `where(cond, x, Invalid)`
but need to keep poisoning.
### 3. `tinygrad/mixin/op.py:280-289` — `_pad_constant` fix
Swapped the `value == 0` early return for `value is Invalid` early return.
When `value is Invalid`, just return `base` (which already has Invalid from
`Ops.PAD`). For all other values (including 0), use the mask approach:
`where(pad(ones_bool), base, const_value)`.
## Current state
- `test/unit/test_invalid_tensor.py`**all 22 pass**
- `test/unit/test_function.py`**5 failures**, all multi-shard tests
## The remaining bug: `cat` + multi-shard
`cat` (op.py:716) uses `pad` + `usum` (element-wise ADD) to combine tensors:
```python
padded = [t.pad(...) for i,t in enumerate(tensors)]
return padded[0].usum(*padded[1:])
```
When two shards are cat'd, each is padded and then summed. The valid masks
are **complementary** (shard 0 valid in positions 0-1, shard 1 valid in 2-3).
`_pad_constant` creates `where(mask_pad, data_pad, 0)` where:
- `mask_pad = where(valid, True, Invalid)` — gate's valid value is const `True`
- `data_pad = where(valid, data, Invalid)` — gate's valid value is loaded `data` (NOT const)
The new const-specific rule handles the mask pad correctly. But for the data pad,
the gate's valid value (`data`) is not a const, so the **non-const** lift-out rule
fires: `where(valid, where(valid, data, Invalid), 0)``where(valid, where(valid, data, 0), Invalid)`.
The `Invalid` else poisons the ADD. The binary Invalid rule lifts both gates out:
`where(c6, data0, Invalid) + where(c8, data1, Invalid)``where(c6&c8, data0+data1, Invalid)`.
Since `c6` and `c8` are complementary, `c6&c8` is always False → result is all Invalid → 0.
### Master comparison
On master, `convert_pad_to_where` uses `0` (not Invalid), so the ADD is just
`where(c6, data0, 0) + where(c8, data1, 0)` with no Invalid, no lifting, works fine.
### Debug output (with changes)
```
c16 = c6.where(c11.index(c13), 0) # where(c6, load0, 0) — correct
c22 = c6.where(0, c17.index(c20)) # where(c6, 0, load1) — correct
c25 = (c6&c8).where((c16+c22), Invalid) # WRONG: c6&c8 always False → all Invalid
```
### Master debug output
```
c13 = c6.where(c8.index(c10), 0) # where(c6, load0, 0)
c21 = c6.where(0, c14.index(c19)) # where(c6, 0, load1)
c22 = c13+c21 # plain ADD, no wrapper — correct
```
## Suggested fix approaches
### Option A: General WHERE simplification rule
Add a rule: `where(a, where(a, x, _), c)``where(a, x, c)`.
When the outer and inner conditions are the same UOp, the inner else is
unreachable. This would simplify `where(valid, where(valid, data, Invalid), 0)`
`where(valid, data, 0)` before the lift-out rule can fire.
Check if this rule already exists in `symbolic.py` — it may need to be added
before the lift-out rules.
### Option B: Don't use Ops.PAD for data in `_pad_constant`
When `value is not Invalid`, avoid creating `Ops.PAD` on the data. Use `cat`
or `expand` to create the padded tensor directly, bypassing the Invalid
propagation entirely.
### Option C: Make the lift-out rule use the outer else value
Change the non-const lift-out rule: when `where(a, where(cond, x, Invalid), c)`
and `c` is a const, use `c` as the else instead of `Invalid`. This is what the
const-specific rule does, but it needs to also handle non-const gate valid values.
## Test commands
```bash
# invalid tensor tests (currently pass)
python -m pytest test/unit/test_invalid_tensor.py -x -q -n12
# function tests (5 multi-shard failures)
python -m pytest test/unit/test_function.py -x -q -n12
# the specific failing test
python -m pytest test/unit/test_function.py::TestFunctionMulti::test_simple_multi_sharded -x -q
# debug the failing case
DEBUG=6 python -c "
from tinygrad import Tensor
a = Tensor([1,2,3,4]).shard(['CPU', 'CPU:1'], axis=0)
print(a.numpy()) # should be [1,2,3,4], gets [0,0,0,0]
"
```
## Lint/typecheck
```bash
python -m mypy tinygrad/
python -m ruff check .
```
+2 -2
View File
@@ -284,8 +284,8 @@ class OpMixin(ElementwiseMixin, ReduceMixin):
X = self.shrink(tuple((-smin(pB,0),smin(pA+s,s)) for (pB,pA),s in zip(pX, self.shape))) if has_neg else self
pads = tuple((smax(pB,0), smax(pA,0)) for pB,pA in pX) if has_neg else pX
base = MovementMixin.pad(X, pads)
if value == 0: return base
if value is not Invalid: base = base.cast(least_upper_dtype(base.dtype, dtypes.from_py(value)))
if value is Invalid: return base
if value != 0: base = base.cast(least_upper_dtype(base.dtype, dtypes.from_py(value)))
return MovementMixin.pad(X.const_like(1).cast(dtypes.bool), pads).where(base, base.const_like(value))
def _pad_circular(self, pX:tuple[tuple[sint, sint], ...]) -> Self:
+2 -2
View File
@@ -1,7 +1,7 @@
from typing import Iterator
import functools, itertools
from dataclasses import dataclass, field, replace
from tinygrad.dtype import dtypes, AddrSpace
from tinygrad.dtype import dtypes, AddrSpace, Invalid
from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, graph_rewrite, sint, AxisType, profile_matches
from tinygrad.uop.ops import consumer_map_from_toposort, gate_kernel_sink
from tinygrad.uop.symbolic import symbolic, pm_simplify_valid, pm_drop_and_clauses
@@ -89,7 +89,7 @@ def convert_pad_to_where_to_keep_behavior_local(ctx:IndexingContext, x:UOp):
if x not in ctx.range_map: return None
bx = create_bufferize_and_index_based_on_ranges(ctx, x)
valid: UOp = UOp.const(dtypes.bool, True).uprod([r.get_valid() for r in ctx.range_map[x][0]])
return valid.where(bx.src[0], UOp.const(x.dtype, 0))
return valid.where(bx.src[0], UOp.const(x.dtype, Invalid))
def convert_reduce_to_reduce_with_ranges(ctx:IndexingContext, x:UOp):
if x.arg[1] == 0: return None
+7
View File
@@ -85,12 +85,19 @@ pm_data_invalid = PatternMatcher([
(UPat(GroupOp.Binary, src=(UPat.var("y"), invalid_gate), name="alu"), lambda cond,x,y,alu,i: cond.where(y.alu(alu.op,x), i.cast(alu.dtype))),
(UPat(GroupOp.Binary-GroupOp.Comparison, src=[invalid_pat, UPat()]), lambda i: i),
# an Invalid condition poisons the whole where; a gated Invalid condition lifts the gate out
# when the gate's valid value is a const (e.g. a pad mask: where(valid, True, Invalid)),
# use the else value in don't-care positions so masks work
(invalid_pat.where(UPat.var("a"), UPat()), lambda i,a: i.cast(a.dtype)),
(UPat.var("cond").where(UPat.cvar("x"), invalid_pat).where(UPat.var("a"), UPat.cvar("b")),
lambda cond,x,i,a,b: cond.where(x.where(a,b), b)),
(invalid_gate.where(UPat.var("a"), UPat.var("b")), lambda cond,x,i,a,b: cond.where(x.where(a,b), i.cast(a.dtype))),
# normalize where(cond, Invalid, val) -> where(~cond, val, Invalid)
(UPat.var("cond").where(invalid_pat, UPat.var("val")), lambda cond, i, val: cond.logical_not().where(val, i) if val.arg != Invalid else i),
# lift Invalid out: a.where(cond.where(x, Invalid), c) -> (~a|cond).where(a.where(x, c), Invalid)
# when a is cond, ~a|cond is True and would drop the Invalid gate (losing the valid), so keep cond as the gate
# when c is a const and the gate's valid value is a const (pad mask), use c in don't-care positions
(UPat.var("a").where(UPat.var("cond").where(UPat.cvar("x"), invalid_pat), UPat.cvar("c")),
lambda cond,i,x,a,c: (cond if a is cond else (a.logical_not()|cond)).where(a.where(x,c), c) if c.arg != Invalid else None),
(UPat.var("a").where(invalid_gate, UPat.var("c")), lambda cond,i,x,a,c:
(cond if a is cond else (a.logical_not()|cond)).where(a.where(x,c), i) if c.arg != Invalid else None),
(UPat.var("a").where(UPat.var("b"), invalid_gate), lambda cond,i,x,a,b: (a|cond).where(a.where(b, x), i) if b.arg != Invalid else None),