Skip to content

Commit 34e3e7a

Browse files
committed
docs: ModelBuildSpec.md — roadmap for Loss/Optim/Rewriters layer
Расширяем IR за пределы forward-цепочки. Сейчас BrickGraph описывает только forward; MTP/IFIM/MHC требуют graph rewrites + custom loss + маршрутизацию градиентов. Делаем ModelBuildSpec = (BrickGraph, LossSpec, OptimSpec, rewrites...) и Rewriter protocol с конкретными MTPRewriter / IFIMRewriter / MHCRewriter. 5 этапов: A — loss_spec.py + optim_spec.py (data-layer) B — ModelBuildSpec + verify_build_spec C — Rewriter protocol + MTPRewriter D — IFIMRewriter + MHCRewriter + composition E — build_model API + executable wiring bd epic: cppmega-mlx-qu2 (5 stage children qu2.1-qu2.5).
1 parent 9919468 commit 34e3e7a

1 file changed

Lines changed: 288 additions & 0 deletions

File tree

ModelBuildSpec.md

Lines changed: 288 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,288 @@
1+
План: ModelBuildSpec Layer (Loss + Optim + Graph Rewrites поверх Bricks)
2+
3+
Цель: расширить IR за пределы forward-цепочки, чтобы MTP / IFIM / MHC
4+
можно было собрать визуально. Сейчас `BrickGraph` описывает только
5+
forward; loss и optim живут руками в `scripts/m04_train_step.py`. Для
6+
MTP это смертельно — он требует переписать сам forward (K параллельных
7+
голов) ПЛЮС multi-loss с per-head весами ПЛЮС маршрутизацию градиентов.
8+
9+
Лежит между BrickGraph и executable runtime; потребляет всё из
10+
`cppmega_v4.spec.*` (ShapeContract, ResolvedBrickGraph, MemoryReport).
11+
12+
---
13+
1. Что уже есть (карта существующего)
14+
15+
Forward IR (cppmega_v4/fusion/brick_graph.py)
16+
BrickGraph (nodes + edges) — топология forward. Знает kind, params,
17+
module. Не знает loss, optim, rewrites.
18+
19+
Spec layer (cppmega_v4/spec/, шипнуто на main, 167 тестов)
20+
- ShapeContract / ResolvedBrickGraph / MemoryReport
21+
- verify_and_estimate(graph) → одношаговая проверка для GUI
22+
- adapters (6 правил с авто-вставкой)
23+
24+
Auto-Fusion (cppmega_v4/fusion/, шипнуто, 175 тестов)
25+
FusionRegionPlan + auto_compile_region + 12 architecture presets.
26+
27+
Тренировочный код (scripts/m04_train_step.py + recipes/)
28+
Hardcoded: один CE loss, AdamW, нет K-head, нет per-head βi/λi.
29+
Если пользователь захочет MTP K=2 — надо ПИСАТЬ КОД, а не натянуть.
30+
31+
Brick `nemotron_h_mtp` (есть в BLOCK_BUILDERS)
32+
Sort-of MTP-блок, но не настоящий K-head rewrite — просто
33+
re-export из mlx-lm PR #1161. Loss-сторону не трогает.
34+
35+
---
36+
2. Что НЕ хватает (gap-analysis)
37+
38+
┌─────────────────────────────────────────┬──────────────────────────────────┐
39+
│ Что нужно GUI │ Сейчас │
40+
├─────────────────────────────────────────┼──────────────────────────────────┤
41+
│ Декларировать loss (CE / MTP / IFIM) │ ✗ (hardcoded в scripts/) │
42+
├─────────────────────────────────────────┼──────────────────────────────────┤
43+
│ Декларировать optimizer + hyperparams │ ✗ (hardcoded — AdamW lr=3e-4) │
44+
├─────────────────────────────────────────┼──────────────────────────────────┤
45+
│ Описать модель целиком (forward+L+O) │ ✗ (только BrickGraph) │
46+
├─────────────────────────────────────────┼──────────────────────────────────┤
47+
│ Применить MTP rewrite к графу │ ✗ (одна голова в forward) │
48+
├─────────────────────────────────────────┼──────────────────────────────────┤
49+
│ Применить IFIM (inverse FIM) shaping │ ✗ │
50+
├─────────────────────────────────────────┼──────────────────────────────────┤
51+
│ Применить MHC (multi-head copy) attn │ ✗ │
52+
├─────────────────────────────────────────┼──────────────────────────────────┤
53+
│ Per-head weights βi / λi для multi-loss │ ✗ │
54+
├─────────────────────────────────────────┼──────────────────────────────────┤
55+
│ Routing градиентов (head_i -> backbone) │ ✗ │
56+
├─────────────────────────────────────────┼──────────────────────────────────┤
57+
│ Проверить coherence loss vs graph │ ✗ │
58+
└─────────────────────────────────────────┴──────────────────────────────────┘
59+
60+
Конкретный gap: на vapor MTP K=2 (см. M0.5 bead `cppmega-mlx-t8f.5`) — у
61+
нас нет ни LossSpec, ни Rewriter — только мечта о том, что m04_train_step
62+
"должен это делать". Этот roadmap закрывает gap визуально-описуемым
63+
ModelBuildSpec.
64+
65+
---
66+
3. Архитектура (новый пакет cppmega_v4/build/)
67+
68+
3.1 cppmega_v4/build/loss_spec.py — декларация loss
69+
70+
class LossKind(str, Enum):
71+
CROSS_ENTROPY = "cross_entropy"
72+
MTP_WEIGHTED = "mtp_weighted" # Σ βi * CE(head_i, shifted_label_i)
73+
IFIM_SHAPED = "ifim_shaped" # CE + λ * IFIM penalty
74+
MHC_ATTN_BIAS = "mhc_attn_bias" # CE + multi-head copy auxiliary
75+
CUSTOM = "custom" # caller supplies callable
76+
77+
@dataclass(frozen=True)
78+
class LossSpec:
79+
kind: LossKind
80+
params: Mapping[str, float] # {"k": 2, "beta_0": 1.0, "beta_1": 0.6, ...}
81+
head_outputs: tuple[str, ...] # names of brick outputs that feed loss
82+
# (must reference nodes that exist in graph
83+
# after rewrites)
84+
label_source: str # "next_token" | "next_k_tokens" | "doc_ids"
85+
reduction: str = "mean"
86+
87+
# Built-ins:
88+
def cross_entropy_loss(head_output_name="logits") -> LossSpec: ...
89+
def mtp_weighted_loss(k=2, beta=(1.0, 0.6), lambda_=0.3) -> LossSpec: ...
90+
def ifim_shaped_loss(lambda_fim=0.1, head_output_name="logits") -> LossSpec: ...
91+
92+
3.2 cppmega_v4/build/optim_spec.py — декларация optimizer
93+
94+
class OptimKind(str, Enum):
95+
ADAMW = "adamw"
96+
MUON = "muon"
97+
MUON_ADAMW_HYBRID = "muon_adamw_hybrid" # 2D params on Muon, rest AdamW
98+
SGD = "sgd"
99+
100+
@dataclass(frozen=True)
101+
class ParamGroup:
102+
"""One param group. The matcher selects parameters; the hyperparams
103+
apply to the matched subset (per-group lr / wd / betas)."""
104+
matcher: str # "all" | "moe_experts" | "embeddings" | regex
105+
lr: float
106+
weight_decay: float = 0.01
107+
betas: tuple[float, float] | None = None # AdamW-only
108+
ns_steps: int | None = None # Muon-only
109+
110+
@dataclass(frozen=True)
111+
class OptimSpec:
112+
kind: OptimKind
113+
groups: tuple[ParamGroup, ...]
114+
gradient_clip_norm: float | None = 1.0
115+
mixed_precision: bool = True
116+
117+
# Built-ins:
118+
def adamw(lr=3e-4, wd=0.01, betas=(0.9, 0.95)) -> OptimSpec: ...
119+
def muon(lr=1e-2, ns_steps=5) -> OptimSpec: ...
120+
def muon_adamw_hybrid(muon_lr=1e-2, adam_lr=3e-4) -> OptimSpec: ...
121+
122+
3.3 cppmega_v4/build/model_build_spec.py — composer
123+
124+
@dataclass(frozen=True)
125+
class ModelBuildSpec:
126+
graph: BrickGraph
127+
loss: LossSpec
128+
optim: OptimSpec
129+
rewrites: tuple["Rewriter", ...] = () # applied in order
130+
dim_env: Mapping[str, int] = field(default_factory=dict)
131+
132+
def apply_rewrites(self) -> "ModelBuildSpec":
133+
"""Run every Rewriter; return new spec (frozen, immutable)."""
134+
...
135+
136+
def verify_build_spec(spec) -> BuildDiagnostics:
137+
"""Check coherence:
138+
- loss.head_outputs ⊆ rewritten_graph.nodes
139+
- optim.groups[*].matcher matches at least one param
140+
- rewrites apply cleanly (no rewrite cycles, no name conflicts)
141+
- shape contracts still valid after rewrites
142+
"""
143+
144+
3.4 cppmega_v4/build/rewriters.py — graph-rewrite слой
145+
146+
class Rewriter(Protocol):
147+
"""A pure graph transformation. Takes a ModelBuildSpec, returns a
148+
new ModelBuildSpec. Must not mutate inputs."""
149+
def __call__(self, spec: ModelBuildSpec) -> ModelBuildSpec: ...
150+
151+
Built-ins:
152+
MTPRewriter(k=2, share_backbone=True)
153+
- Find the head brick (usually `logits` projection)
154+
- Materialise K copies, each predicting shifted labels [+0, +1, ..., +k-1]
155+
- Adjust LossSpec: convert single-output CE into mtp_weighted
156+
(if loss is CE; otherwise raise)
157+
- Update graph: K new head nodes; backbone shared (share_backbone=True)
158+
- Update OptimSpec: optional separate param group for new heads
159+
160+
IFIMRewriter(lambda_fim=0.1)
161+
- Wrap the final head's output in an IFIM auxiliary loss node
162+
- Adds a `ifim_aux` virtual brick that consumes logits + computes
163+
the Fisher-information penalty
164+
- LossSpec rewritten to ifim_shaped
165+
166+
MHCRewriter(num_copies=2)
167+
- For each attention brick, materialise N copies sharing weights
168+
- Add a "multi-head copy" auxiliary loss
169+
170+
3.5 cppmega_v4/build/api.py — public
171+
172+
def build_model(spec: ModelBuildSpec) -> BuiltModel:
173+
"""Apply rewrites, verify, materialise nn.Module + loss callable +
174+
optimizer instance."""
175+
176+
@dataclass(frozen=True)
177+
class BuiltModel:
178+
module: nn.Module # post-rewrite forward
179+
loss_fn: Callable[..., mx.array]
180+
optimizer: object # MLX-native optim instance
181+
param_groups: tuple[tuple[str, ...], ...] # group_idx -> param names
182+
spec_applied: ModelBuildSpec # post-rewrite snapshot for telemetry
183+
184+
def verify_build_spec(spec) -> BuildDiagnostics:
185+
"""Shape-coherence + loss-head-output coverage + optim-matcher coverage."""
186+
187+
---
188+
4. Какие rewrites нужны для текущих presets
189+
190+
┌─────────────────────────┬───────────────────────────────────────────────┐
191+
│ Preset │ Recommended rewrites │
192+
├─────────────────────────┼───────────────────────────────────────────────┤
193+
│ qwen3_next │ MTPRewriter(k=2, beta=(1.0,0.6)) │
194+
│ kimi_linear / kimi_k2 │ MTPRewriter(k=2) │
195+
│ deepseek_v3 / v4_flash │ MTPRewriter(k=2) — DeepSeek уже использует │
196+
│ gemma4 │ MTPRewriter(k=3) — Gemma-4 drafter │
197+
│ mistral4 │ — (нет MTP) │
198+
│ ling26 │ MTPRewriter(k=2) + IFIMRewriter(λ=0.05) │
199+
│ longcat │ — │
200+
│ nemotron3 │ замена `nemotron_h_mtp` brick на MTPRewriter │
201+
│ zaya1 │ MTPRewriter(k=2) + MHCRewriter(num_copies=2) │
202+
│ arcee_trinity │ MTPRewriter(k=2) │
203+
└─────────────────────────┴───────────────────────────────────────────────┘
204+
205+
---
206+
5. Что GUI получает
207+
208+
- В правой панели: «Loss» dropdown + per-head βi sliders + checkbox для
209+
IFIM / MHC. Сразу видно когда βi не складываются в 1.0 (warning).
210+
- В правой панели: «Optimizer» dropdown (AdamW / Muon / Hybrid) + lr
211+
slider; per-group панель когда у юзера несколько param-groups (MoE
212+
experts vs main).
213+
- Над графом: цепочка «Rewrites applied» (chips) — кликабельные
214+
MTPRewriter(k=2), IFIMRewriter(λ=0.05). Юзер может drag-drop порядок,
215+
выкинуть, добавить.
216+
- После применения rewrites — preview графа с подсвеченными новыми
217+
нодами (head_1, head_2, ifim_aux) в другом цвете.
218+
- Memory report пересчитывается на post-rewrite graph — юзер видит
219+
«MTP K=2 adds 1.8 GB» прежде чем нажмёт Train.
220+
221+
---
222+
6. План реализации поэтапно
223+
224+
Этап A — LossSpec + OptimSpec (1 заход)
225+
- cppmega_v4/build/loss_spec.py — LossKind, LossSpec, built-ins
226+
- cppmega_v4/build/optim_spec.py — OptimKind, ParamGroup, OptimSpec, built-ins
227+
- Чистый data-layer, никакой MLX runtime; всё валидируется в __post_init__.
228+
- Тесты: каждый built-in возвращает валидный spec; rejection-тесты
229+
на негативные lr, плохой kind, пустые head_outputs.
230+
231+
Этап B — ModelBuildSpec + verify_build_spec (1 заход)
232+
- cppmega_v4/build/model_build_spec.py — composer (immutable dataclass)
233+
- apply_rewrites — chain order semantics
234+
- BuildDiagnostics + verify_build_spec (shape-coherence через
235+
cppmega_v4.spec, loss head_outputs ⊆ graph.names, optim matcher
236+
coverage)
237+
- Тесты: верифицирует Qwen3-Next + AdamW + CE — clean; искусственный
238+
bad head_output — ERROR; пустой matcher — WARNING.
239+
240+
Этап C — Rewriter protocol + MTPRewriter (1 заход)
241+
- cppmega_v4/build/rewriters.py — Rewriter Protocol + MTPRewriter
242+
- MTPRewriter материализует K head копий, перепишет LossSpec
243+
CE → mtp_weighted, добавит head-only param group в OptimSpec
244+
- Тесты: K=1 — no-op; K=2 — 2 head nodes + mtp_weighted loss; K=3 +
245+
share_backbone=False — duplicate-backbone path; ошибка когда loss
246+
не CE (нельзя автоматически переписать)
247+
- System: применить к qwen3_next preset, verify — clean, memory rate
248+
растёт ровно на N(heads-1) * params(head).
249+
250+
Этап D — IFIMRewriter + MHCRewriter (1 заход)
251+
- cppmega_v4/build/rewriters.py — добавить два rewriter
252+
- IFIMRewriter — добавит aux node + λ-loss
253+
- MHCRewriter — на каждом attention брике сделает copies
254+
- Тесты: composition (MTPRewriter + IFIMRewriter), порядок
255+
применения, anti-cycle проверка.
256+
257+
Этап E — build_model + executable wiring + GUI integration test (1 заход)
258+
- cppmega_v4/build/api.py — build_model(spec) → BuiltModel
259+
- Wire BrickGraph nodes (с module) в nn.Module-композицию
260+
- Wire LossSpec → callable (CE / mtp_weighted / IFIM)
261+
- Wire OptimSpec → mlx.optimizers.AdamW / Muon / Hybrid
262+
- Perf-критерий: build_model для каждого из 12 presets < 200 ms
263+
- Тесты: построить qwen3_next + MTPRewriter(k=2) + AdamW → forward
264+
проходит, loss скаляр финитен, optimizer.update без crash, GIU
265+
workflow integration (build_spec → verify → build_model → step).
266+
267+
---
268+
7. Бюджет и риски
269+
270+
Бюджет: 5 этапов × ~2 часа = ~10 часов чистого кодинга. A и B —
271+
data-layer, последовательны. C/D — rewriter implementations, частично
272+
параллелятся (MTP vs IFIM/MHC независимы). E — wiring + integration.
273+
274+
Главный риск: MTPRewriter должен корректно "найти head brick" в графе.
275+
Сейчас наш `logits` head — это часть UnifiedSuperblockV4, не отдельный
276+
brick. Mitigation: ввести опциональный marker `is_head=True` в BrickNode
277+
или соглашение «head — последний node в графе». Документировать
278+
жёстко.
279+
280+
Второй риск: LossSpec coherence vs Rewriter. Если юзер ставит CE,
281+
потом MTPRewriter автоматически перепишет в mtp_weighted — это сюрприз.
282+
Mitigation: verify_build_spec показывает что **итоговая** loss будет
283+
mtp_weighted, ДО build_model.
284+
285+
Третий риск: composition rewrites в неправильном порядке. IFIMRewriter
286+
после MTPRewriter работает на K голов или на одну? Mitigation:
287+
Rewriter имеет required_precondition + provided_postcondition; verifier
288+
выкинет ERROR если порядок неверен.

0 commit comments

Comments
 (0)