|
| 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