"""Phase 1 spec test for ``tl.copy_to`` (ADR-0063 §D3.1). ADR-0063 §D3.1: ``tl.copy_to(dst, src)`` is a TCM-to-TCM byte copy that lets a kernel write a scoped result's bytes to a persistent (outside- scope) address. Required for the two-arena flash pattern in §D3. Emit-time validation (ADR-0063 §D3.1 "Mechanics" / "Emit-time validation"): - ``dst.shape == src.shape`` - ``dst.dtype == src.dtype`` - ``dst.space == "tcm"`` (TCM-only — HBM goes through ``tl.store``) - ``src.space == "tcm"`` (same reason) op_log shape (ADR-0063 §D3.1 "Mechanics"): - ``op_kind="math"`` (runs on the vector engine) - ``op_name="copy"`` Phase 1 (this commit): tests only — production code lands in Phase 2. The tests are written against the ``CopyCmd`` dataclass and the ``TLContext.copy_to`` method that ADR-0063 §D3.1 declares; both are absent today, so every test in this file should currently fail with ``AttributeError`` / ``ImportError``. """ from __future__ import annotations import pytest from kernbench.triton_emu.tl_context import TLContext def _ctx() -> TLContext: """TLContext with a real scratch base so _make_compute_out allocates.""" return TLContext( pe_id=0, num_programs=1, dispatch_cycles=0, scratch_base=1 << 24, scratch_size=1 << 20, ) # ── 1. Method exists on the tl surface ─────────────────────────────── def test_copy_to_method_exists(): """ADR-0063 §D3.1: TLContext must expose ``copy_to(dst, src)``.""" ctx = _ctx() assert hasattr(ctx, "copy_to"), ( "TLContext must expose copy_to(dst, src) per ADR-0063 §D3.1" ) # ── 2. Emits a CopyCmd with the documented fields ──────────────────── def test_copy_to_emits_copy_command(): """ADR-0063 §D3.1 Mechanics: a new ``CopyCmd(src, dst, nbytes)`` command is recorded; data_op=True (so it appears in op_log).""" from kernbench.common.pe_commands import CopyCmd # added in Phase 2 ctx = _ctx() src = ctx._make_compute_out(shape=(8, 64), dtype="f16") dst = ctx._make_compute_out(shape=(8, 64), dtype="f16") ctx.copy_to(dst, src) copy_cmds = [c for c in ctx.commands if isinstance(c, CopyCmd)] assert len(copy_cmds) == 1, ( f"exactly one CopyCmd expected; got {len(copy_cmds)}" ) cmd = copy_cmds[0] assert cmd.src is src assert cmd.dst is dst assert cmd.nbytes == 8 * 64 * 2 assert cmd.data_op is True # ── 3. Shape mismatch is rejected at emit time ─────────────────────── def test_copy_to_shape_mismatch_rejected(): """ADR-0063 §D3.1: ``dst.shape == src.shape`` enforced at emit time.""" ctx = _ctx() src = ctx._make_compute_out(shape=(8, 64), dtype="f16") dst = ctx._make_compute_out(shape=(8, 128), dtype="f16") # mismatched with pytest.raises((ValueError, AssertionError)) as excinfo: ctx.copy_to(dst, src) assert "shape" in str(excinfo.value).lower() # ── 4. Dtype mismatch is rejected at emit time ─────────────────────── def test_copy_to_dtype_mismatch_rejected(): """ADR-0063 §D3.1: ``dst.dtype == src.dtype`` enforced at emit time.""" ctx = _ctx() src = ctx._make_compute_out(shape=(8, 64), dtype="f16") dst = ctx._make_compute_out(shape=(8, 64), dtype="f32") # mismatched with pytest.raises((ValueError, AssertionError)) as excinfo: ctx.copy_to(dst, src) assert "dtype" in str(excinfo.value).lower() # ── 5. Non-TCM dst is rejected at emit time ────────────────────────── def test_copy_to_dst_must_be_tcm(): """ADR-0063 §D3.1: ``dst.space == 'tcm'`` enforced at emit time. Writing to HBM goes through ``tl.store``, not ``tl.copy_to`` — keeping copy_to TCM-only avoids polluting op_log with DMA entries that the bump-allocator scope mechanism doesn't model. """ ctx = _ctx() src = ctx._make_compute_out(shape=(8, 64), dtype="f16") # Fabricate an HBM-resident dst (e.g. a load handle). dst_hbm = ctx.load(0x10_000, shape=(8, 64), dtype="f16") assert dst_hbm.space == "hbm", "fixture sanity check" with pytest.raises((ValueError, AssertionError)) as excinfo: ctx.copy_to(dst_hbm, src) msg = str(excinfo.value).lower() assert "tcm" in msg or "space" in msg or "hbm" in msg # ── 6. Non-TCM src is rejected at emit time ────────────────────────── def test_copy_to_src_must_be_tcm(): """ADR-0063 §D3.1: src.space == 'tcm' (symmetric to dst). Reading from HBM goes through ``tl.load``, not ``tl.copy_to``. """ ctx = _ctx() src_hbm = ctx.load(0x10_000, shape=(8, 64), dtype="f16") assert src_hbm.space == "hbm" dst = ctx._make_compute_out(shape=(8, 64), dtype="f16") with pytest.raises((ValueError, AssertionError)) as excinfo: ctx.copy_to(dst, src_hbm) msg = str(excinfo.value).lower() assert "tcm" in msg or "space" in msg or "hbm" in msg # ── 7. op_log routes CopyCmd as ("math", "copy", ...) ──────────────── def test_copy_to_op_log_classifies_as_math_copy(): """ADR-0063 §D3.1 Mechanics: op_kind='math', op_name='copy'. The vector engine handles the copy; op_log classification reflects that (engine-class accounting in bench summaries can sum over op_kind='math' to include copy time alongside softmax/exp). """ from kernbench.common.pe_commands import CopyCmd # added in Phase 2 from kernbench.sim_engine.op_log import _extract_op_info ctx = _ctx() src = ctx._make_compute_out(shape=(8, 64), dtype="f16") dst = ctx._make_compute_out(shape=(8, 64), dtype="f16") ctx.copy_to(dst, src) cmd = next(c for c in ctx.commands if isinstance(c, CopyCmd)) op_kind, op_name, params = _extract_op_info(cmd) assert op_kind == "math", ( f"copy runs on the vector engine: op_kind='math'; got {op_kind!r}" ) assert op_name == "copy", ( f"op_name must be 'copy' per ADR-0063 §D3.1; got {op_name!r}" ) # Params must let DataExecutor replay (read src, write dst). assert params.get("dst_addr") == dst.addr assert params.get("dst_space", "tcm") == "tcm"