File size: 2,399 Bytes
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
# SPDX-License-Identifier: Apache-2.0
"""tt-nn implementation of Diffusion Planner v5.0 for one Blackhole p150 (12x10 grid, ETH dispatch).

Planned layout (PORT_LOG.md milestones M1-M5; nothing here yet at M0):

- ``model.py``     ``TtDiffusionPlanner(device, weights, knobs=None, debug=False)``: the canonical weights of
                   ``reference.weights`` -> device tensors (bf16; the ego / neighbour pre-projection island in fp32,
                   pad-relative with the constants of ``reference.rewrites.island_constants``; the per-step adaLN tables
                   of ``reference.rewrites.adaln_tables`` folded into LayerNorm affine rows; the solver coefficients of
                   ``host.solver.solver_plan``), and the variants of a ``ttaw.trace.TraceRunner``: ``plan`` (encoder ->
                   fusion -> hoisted cross K|V -> 11 x (DiT + fp32 DPM-Solver++(2M) update + prefix constraint) -> turn
                   head, one packed readback), plus the debug variants ``encoder_taps`` and ``decode_once`` used by
                   ``tests/test_pcc_device.py``. ``forward(prepared) -> {"final_x0", "logit"[, "denoising_steps"]}``.
- ``encoder.py``   MLP-Mixer trunk (token mixing as transpose -> 2-D ``[1, 1, E*128, K] @ W`` -> transpose, probe P12),
                   entity heads, positional embedding, the fusion transformer (C20 SDPA 564 x 564 with the -inf key
                   mask).
- ``decoder.py``   one DiT evaluation with table-driven modulation, masked self-attention (321 x 321) and
                   cross-attention on the hoisted K|V (321 x 564); the solver update and prefix constraint.
- ``kernels/``     none planned for the functional port (PLAN.md 2.12: no custom kernels); fused mixer / DiT-step
                   kernels and the single-megakernel attempt are optimization work (PLAN.md 5, D19).

Rules: read the grid from the device (``ttaw.device.compute_grid``); explicit ``compute_kernel_config`` on every
matmul (``ttaw.precision``; HiFi4 for fp32 operands); ``epsilon=1e-5`` on every LayerNorm; ``is_causal=False`` and
``scale=1/sqrt(32)`` on every SDPA (through the C20 wrapper ``ttaw.ops.attention``); no host reads inside a capture;
no program compiled after the first capture. ``api.py`` imports this package inside ``_build`` only, so
``import tt_diffusion_planner`` stays free of ttnn while the modules here may import it at the top.
"""