Skip to content

fix(analysis): count an invariant operand's re-reads in traffic total #168

Description

@zhen8838

Why

Spread states the two traffic accounts a report carries
(docs/spec/analysis.md:234-241):

logical: What the authored operation asks for before replication.
total:   What the whole execution asks for.

_shares fills both from the same entry (src/tilefoundry/analysis/memory.py:278-281):

Spread(
    logical=shares.get(_WHOLE, TrafficBytes()),
    total=shares.get(_WHOLE, TrafficBytes()),
    per_unit=...,
)

and the entry they share counts an occurrence only over the scopes it varies in
(memory.py:395-401):

repeats = 1
cursor = ctx.current
while cursor.parent is not None:
    if cursor.is_variant(expr):
        repeats *= max(1, cursor.trips())
    cursor = cursor.parent
add_traffic(ctx.totals, ctx.shares, moved, repeats)

A tiled GEMM reads A once per N tile and B once per M tile. Neither re-read is
counted, because A is invariant in n and B is invariant in m. So total
reports the bytes the program touches, which is what logical already reports,
and no field reports what the execution moves.

Two tilings of one 256x256x256 bf16 GEMM, differing only in BM/BN (program
below, tilefoundry analyze prog.py:TILED out --memory --json, at f538e08b):

BM x BN reported gmem read (logical = total) bytes the nest reads roofline.memory_ns
64 x 64 262144 1048576 110
128 x 128 262144 524288 110

One tiling moves twice the data of the other and the report does not distinguish
them. roofline builds memory_ns on this account, so bound-by compares
compute against a floor rather than against the traffic.

The case this came from is a bf16 GEMM at M=8192, K=5120, N=17408. Four tilings --
128x64, 64x128, 256x64, 64x256 -- all report gmem: r262144000, while their nests
read 31.88, 31.88, 26.56 and 26.56 GiB. Choosing a tile shape is choosing the reuse,
and the analysis that is supposed to price that choice returns one number for all
four.

What

Give total the meaning the spec already assigns it, and leave logical alone:

  • memory.py:395-401 keep the variant-only product for the logical account and
    multiply by every enclosing scope's trips for the total account.
  • memory.py:278-281 stop filling both fields from shares[_WHOLE].
  • roofline read total for memory_ns, so a compute/memory verdict is stated
    against the bytes that move.

Whether a cache model belongs between the two accounts is a separate question. The
ratio between them is the reuse a tiling assumes; reporting both states it without
deciding it.

Contract

docs/spec/analysis.md:234-241 already states what the two fields mean, so this is
the implementation meeting the spec rather than a new contract. The traffic section
of the report format does not say which account a reader is looking at; if the two
now differ, it should.

Risk

compute_cost reads the same logical/total pair (compute_cost.py:218-221), so
the two analyses would keep one reading. Changing total moves roofline.memory_ns
and therefore some bound-by verdicts: a program reported compute-bound against the
floor can become memory-bound against the moved bytes. That is the correction, but it
will move existing report expectations.

Repro program (only BM/BN differ between the two runs)
from tilefoundry import func, module
from tilefoundry.dsl import Mesh, Tensor, Topology, tf
from tilefoundry.target import CudaTarget

M = K = N = 256
BM = BN = 64          # second run: 128
BK = 32


@module(entry="gemm", target=CudaTarget("nvidia.h200_sxm"),
        topologies=(Topology("cta", 1), Topology("thread", 128)))
class TILED:
    @func
    def gemm(a: Tensor[(M, K), "bf16"],
             b: Tensor[(K, N), "bf16"]) -> Tensor[(M, N), "bf16"]:
        with Mesh(("cta",), layout=(1,), names=("g",)) as cta:
            result = tf.zeros(Tensor[(M, N), "bf16"])
            for m in tile(M, BM):
                for n in tile(N, BN):
                    acc = tf.zeros(Tensor[(BM, BN), "f32", (BM, BN), "rmem"])
                    for k in tile(K, BK):
                        a_s = tf.reshard(a[m, k], (BM, BK), "smem")
                        b_s = tf.reshard(b[k, n], (BK, BN), "smem")
                        acc = acc + tf.reshard(
                            tf.matmul(tf.cast(a_s, dtype="f32"), tf.cast(b_s, dtype="f32")),
                            (BM, BN), "rmem")
                    result = tf.insert_slice(result, tf.cast(acc, dtype="bf16"), (m, n))
            return result

Per-value annotations show where the factor goes: at BM=BN=64 the A reshard is
gmem:r4096 and is counted 32 times (m 4 x k 8), giving 131072 = A once. The
n loop's 4 trips are not in the product.

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