tl.composite: fused epilogue ops with per-op scope
Extend tl.composite() with an ordered epilogue list. Each op carries
a scope flag - output_tile (default, runs once per (m,n) before
STORE), k_tile (every K-tile right after GEMM), or kernel. Plan
generator slots MATH stages by scope; pe_math reuses pe_dma's
local-loop pattern so chained epilogues (bias->relu) skip the port
hop. op_log captures per-stage params for telemetry. Topology
gains a gemm->math edge (snapshot test updated).
API stays backward-compatible - `epilogue=` is opt-in.
Example:
h = tl.composite(
op="gemm", a=a, b=b, out_ptr=int(out),
epilogue=[
{"op": "dequant", "scale": s_per_k, "scope": "k_tile"},
{"op": "bias", "bias": bias_vec},
{"op": "relu"},
{"op": "scale", "factor": 0.5},
],
)
tl.wait(h)
Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -157,6 +157,9 @@ class PeSchedulerComponent(ComponentBase):
|
||||
b = cmd.b
|
||||
M, K = a.shape[-2], a.shape[-1]
|
||||
N = b.shape[-1]
|
||||
# When CompositeCmd.ops is populated, ops[0] is the head and
|
||||
# ops[1:] is the epilogue spec list. Empty ops → legacy path.
|
||||
epi_specs = tuple(cmd.ops[1:]) if cmd.ops else ()
|
||||
return generate_gemm_plan(
|
||||
M=M, K=K, N=N,
|
||||
tile_m=self.TILE_M, tile_k=self.TILE_K, tile_n=self.TILE_N,
|
||||
@@ -165,6 +168,7 @@ class PeSchedulerComponent(ComponentBase):
|
||||
pe_prefix=pp,
|
||||
a_pinned=getattr(a, "pinned", False),
|
||||
b_pinned=getattr(b, "pinned", False),
|
||||
epilogue_specs=epi_specs,
|
||||
)
|
||||
else:
|
||||
# Math composite
|
||||
|
||||
Reference in New Issue
Block a user