Why
check cannot establish a reference value for any program that reads its mesh
coordinate. The evaluator models a single mesh participant by design
(docs/spec/evaluator.md:225: Reshard "performs no cross-participant data
movement", Local "returns its operand's value for the single modelled
participant"), so m.w has no value:
m.w desugars to: Arange((4,)) -> Reshard(S(0)) -> Local() -> Reshape(())
intended: [0,1,2,3] split by mesh this unit -> (1,) -> scalar u
today: [0,1,2,3] identity identity -> (4,) -> crash
$ tilefoundry check <fixture>:Strided --inputs random --out output --fn nan_inf
tilefoundry check: error: shape '[]' is invalid for input of size 4
This blocks the ordinary way to split a long reduction over a worker axis
(b0 = t + m.w * BLK) — the only placement with enough parallel units for
long-context attention. tests/fixtures/placed/flash_split_k_decode.py:59
already has this shape.
What
Two facts worth recording before anyone picks this up.
The unsound set is exactly one condition. parser/pattern_nodes.py:2903 is
the only site that constructs Local, so reading the coordinate is the only
source of per-unit variation. When nothing varies per unit, every unit computes
the same value and today's identity reading coincides with the truth — that is
why Fixed-shaped programs pass. A static marking
unit_varying(e) = e is Local on a Split axis | any operand is unit_varying
measured over all of tests/fixtures/ — 54 functions, 814 Call nodes — marks
55 nodes (6.8%), all inside the two functions that read the coordinate.
Everything else is unaffected by whatever fix lands.
Neither JAX nor Shardy has a reference evaluator to copy. Shardy is
representation + propagation + SPMD partitioner; execution is delegated to the
backend. shard_map's reference semantics is
concatenate([f(blk) for blk in split(y)]) with per-device local shapes, and
that equation only holds for a collective-free body — ours has one
(Reshard S→B is an all-gather).
Contract
docs/spec/evaluator.md §6 has to change: today's three statements
(single participant / Reshard does not move data / Local is identity) are
what any fix replaces. Until then, check's answer for a mesh program is only
correct when no node reads the coordinate.
Risk
- Refusing this path is already going out (
fix/failure-diagnostics), so the
failure is diagnosable in the meantime; this issue is about computing it.
docs/spec/evaluator.md does not say what TensorValue.data holds for a
Partial value. Today Reshard P→B is identity, which is right under
"data is already reduced" and wrong under "data is one unit's partial".
No fixture produces it, but ir/hir/tensor/index_select.py:78 and
visitor_registry/shard_propagate.py can. Settle this before implementing.
- Separately, typeinfer accepts an op reducing across a
Split axis with no
intervening Reshard (probe builds cleanly). The evaluator's answer there is
right; the legality question is not this issue's.
Why
checkcannot establish a reference value for any program that reads its meshcoordinate. The evaluator models a single mesh participant by design
(
docs/spec/evaluator.md:225:Reshard"performs no cross-participant datamovement",
Local"returns its operand's value for the single modelledparticipant"), so
m.whas no value:This blocks the ordinary way to split a long reduction over a worker axis
(
b0 = t + m.w * BLK) — the only placement with enough parallel units forlong-context attention.
tests/fixtures/placed/flash_split_k_decode.py:59already has this shape.
What
Two facts worth recording before anyone picks this up.
The unsound set is exactly one condition.
parser/pattern_nodes.py:2903isthe only site that constructs
Local, so reading the coordinate is the onlysource of per-unit variation. When nothing varies per unit, every unit computes
the same value and today's identity reading coincides with the truth — that is
why
Fixed-shaped programs pass. A static markingmeasured over all of
tests/fixtures/— 54 functions, 814Callnodes — marks55 nodes (6.8%), all inside the two functions that read the coordinate.
Everything else is unaffected by whatever fix lands.
Neither JAX nor Shardy has a reference evaluator to copy. Shardy is
representation + propagation + SPMD partitioner; execution is delegated to the
backend.
shard_map's reference semantics isconcatenate([f(blk) for blk in split(y)])with per-device local shapes, andthat equation only holds for a collective-free body — ours has one
(
ReshardS→B is an all-gather).Contract
docs/spec/evaluator.md§6 has to change: today's three statements(single participant /
Resharddoes not move data /Localis identity) arewhat any fix replaces. Until then,
check's answer for a mesh program is onlycorrect when no node reads the coordinate.
Risk
fix/failure-diagnostics), so thefailure is diagnosable in the meantime; this issue is about computing it.
docs/spec/evaluator.mddoes not say whatTensorValue.dataholds for aPartialvalue. TodayReshardP→B is identity, which is right under"data is already reduced" and wrong under "data is one unit's partial".
No fixture produces it, but
ir/hir/tensor/index_select.py:78andvisitor_registry/shard_propagate.pycan. Settle this before implementing.Splitaxis with nointervening
Reshard(probe builds cleanly). The evaluator's answer there isright; the legality question is not this issue's.