Skip to content

feat(evaluator): model the mesh so check can reference a coordinate-reading program #138

Description

@zhen8838

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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions