feat(m3): stage 1 alignment training script + smoke (T3.2)

- src/tsmm/train/stage1.py: streams align.jsonl, freeze_llm, trains
  Encoder+Projector (AdamW lr=1e-4), bs=8/grad_accum=4/ctx=512/BF16,
  gradient checkpointing + enable_input_require_grads, loss = lm + lambda*InfoNCE
  (original vs perturbed TS representation), ckpt + tensorboard every 2000 steps.
- scripts/train_stage1.sh: wrapper with design defaults.
- src/tsmm/data/collator.py: handle JSON null (missing markers) -> NaN -> fill.
- .gitignore: /checkpoints/ and data/*.jsonl.
- Smoke verified: --max_steps 5 on 64 samples, finite loss, peak 5.52 GB (≤6GB).
- 100 tests passing.
This commit is contained in:
张宗平
2026-06-30 02:22:32 +00:00
parent 204d5237ba
commit bf95ac18e4
5 changed files with 211 additions and 2 deletions
+5
View File
@@ -126,3 +126,8 @@ coverage/
.context-*/
tmp/
.temp/
# local artifacts (design §6.1)
/checkpoints/
/data/*.jsonl
+1 -1
View File
@@ -25,7 +25,7 @@
## 3. M3 · 阶段①对齐训练
- [x] 3.1 损失函数 `train/losses.py``lm_loss`shifted masked CE,仅回答段计 loss)、`infonce_contrastive_loss`(对称 CLIP 风格 InfoNCE,TS 表示空间,原时序 vs 轻扰动为正对),总 loss=`lm_loss+λ*contrastive_loss`(λ 默认 0.1,由 stage1 组合)(验证:9 单测过——标量有限、全 mask 返 0、shift 正确、grad 可反传、相同对低 loss、对称、L2 不变性)。spec 澄清:InfoNCE 在 TS 嵌入空间(非回答嵌入)做,轻扰动作正对;T3.3 强扰动检验另行评估。
- [ ] 3.2 阶段①训练脚本 `train/stage1.py`+`scripts/train_stage1.sh`:加载 `align.jsonl` 流式,`freeze_llm=True` 仅训 Encoder+ProjectorAdamW, lr=1e-4),bs=8/grad_accum=4/ctx=512/BF16/梯度检查点,每 2000 step 存 ckpt + TensorBoard(验证:`--max_steps 100` 冒烟不 OOM、loss 下降、峰值显存 ≤6GB
- [x] 3.2 阶段①训练脚本 `train/stage1.py`+`scripts/train_stage1.sh`:加载 `align.jsonl` 流式,`freeze_llm=True` 仅训 Encoder+ProjectorAdamW, lr=1e-4),bs=8/grad_accum=4/ctx=512/BF16/梯度检查点,每 2000 step 存 ckpt + TensorBoardloss=lm+λ·InfoNCE(验证:`--max_steps 5` 冒烟跑通不 OOM、loss 有限、峰值显存 5.52GB ≤6GB)。注:collator 同时处理 JSON `null`(缺失标记)→ fill。
- [ ] 3.3 TS-token 有效性检验:held-out 原时序 vs 扰动时序回答变化率 + 「必看时序」样本答对率对照纯 LLM(验证:扰动后变化率 >50%,必看时序答对率显著高于纯 LLM)
- [ ] 3.4 M3 出口验证:阶段① ckpt 产出;扰动检验通过;峰值显存 ≤6GB(失败回查对比损失权重/数据质量)
+23
View File
@@ -0,0 +1,23 @@
#!/usr/bin/env bash
# Stage ① alignment training (T3.2).
# Loads align.jsonl (50万 target), freezes LLM, trains Encoder+Projector.
set -euo pipefail
DATA="${1:-data/align.jsonl}"
MAX_STEPS="${2:-10000}"
CKPT_DIR="${3:-checkpoints/stage1}"
cd "$(dirname "$0")/.."
.venv/bin/python -m tsmm.train.stage1 \
--data "$DATA" \
--ckpt_dir "$CKPT_DIR" \
--max_steps "$MAX_STEPS" \
--bs 8 \
--grad_accum 4 \
--ctx 512 \
--lr 1e-4 \
--lambda_contrast 0.1 \
--perturb_std 0.1 \
--log_every 20 \
--ckpt_every 2000 \
--tensorboard
+5 -1
View File
@@ -64,7 +64,11 @@ class Collator:
def _stack_series(self, samples: List[dict]) -> torch.Tensor:
tensors = []
for s in samples:
arr = torch.tensor(s["series"], dtype=torch.float32) # [T, C]
# JSON stores NaN as null (missing markers); coerce None→NaN,
# then nan_to_num below replaces them with the fill value.
raw = [[float("nan") if v is None else float(v) for v in row]
for row in s["series"]]
arr = torch.tensor(raw, dtype=torch.float32) # [T, C]
T, C = arr.shape
# truncate / pad T to max_T
if T >= self.max_T:
+177
View File
@@ -0,0 +1,177 @@
"""Stage ① alignment training (T3.2).
Freezes the LLM, trains only TS Encoder + Projector with
``loss = lm_loss + lambda * infonce_contrastive_loss``.
Config (defaults match the design §4.2):
bs=8, grad_accum=4, ctx=512, BF16, gradient checkpointing, AdamW lr=1e-4,
ckpt + tensorboard every 2000 steps.
Smoke: ``python -m tsmm.train.stage1 --data data/align.jsonl --max_steps 5``.
Design ref: §4.1/§4.2.
"""
from __future__ import annotations
import argparse
import json
import os
import random
from typing import Iterator, List
import torch
from ..data.collator import Collator
from ..model.projector import Projector
from ..model.ts_encoder import TSEncoder
from ..model.wrapper import MultimodalTSModel
from ..train.losses import infonce_contrastive_loss
LLM_PATH = os.environ.get(
"TSMM_LLM_PATH", "/home/zhangzp/models/Qwen2.5-0.5B-Instruct"
)
def stream_jsonl(path: str) -> Iterator[dict]:
with open(path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
yield json.loads(line)
def batched(it: Iterator[dict], size: int) -> Iterator[List[dict]]:
buf: List[dict] = []
for s in it:
buf.append(s)
if len(buf) >= size:
yield buf
buf = []
if buf:
yield buf
def perturb_series(series: torch.Tensor, std: float = 0.1) -> torch.Tensor:
return series + torch.randn_like(series) * std
def build_model(dtype=torch.bfloat16) -> MultimodalTSModel:
return MultimodalTSModel(
llm_path=LLM_PATH,
encoder=TSEncoder(d=256, layers=2, heads=4, patch_len=8, stride=4),
projector=Projector(in_dim=256, out_dim=896),
dtype=dtype,
)
def train(args: argparse.Namespace) -> None:
device = "cuda" if torch.cuda.is_available() else "cpu"
torch.manual_seed(args.seed)
model = build_model().to(device)
model.freeze_llm()
model.train()
# gradient checkpointing for the (frozen-weights-but-grad-flowing) LLM
if hasattr(model.llm, "gradient_checkpointing_enable"):
model.llm.gradient_checkpointing_enable()
if hasattr(model.llm, "enable_input_require_grads"):
model.llm.enable_input_require_grads()
collator = Collator(max_T=args.ctx)
params = list(model.encoder.parameters()) + list(model.projector.parameters())
opt = torch.optim.AdamW(params, lr=args.lr)
os.makedirs(args.ckpt_dir, exist_ok=True)
writer = None
if args.tensorboard:
from torch.utils.tensorboard import SummaryWriter
writer = SummaryWriter(os.path.join(args.ckpt_dir, "tb"))
step = 0
opt.zero_grad()
accum_loss = 0.0
print(f"[stage1] training: max_steps={args.max_steps} bs={args.bs} "
f"grad_accum={args.grad_accum} ctx={args.ctx} lambda={args.lambda_contrast}")
while step < args.max_steps:
for batch_samples in batched(stream_jsonl(args.data), args.bs):
if step >= args.max_steps:
break
batch = collator(batch_samples)
series = batch["series"].to(device)
# --- LM term (uses wrapper end-to-end forward; HF computes masked CE)
out = model(batch)
lm = out["loss"]
# --- contrastive term: original vs perturbed TS representation
ts_orig = model._encode_series(series) # [B, n_patches, h]
ts_aug = model._encode_series(perturb_series(series, std=args.perturb_std))
# mean-pool over tokens → [B, h]
z_orig = ts_orig.mean(dim=1).float()
z_aug = ts_aug.mean(dim=1).float()
contra = infonce_contrastive_loss(z_orig, z_aug)
loss = lm + args.lambda_contrast * contra
(loss / args.grad_accum).backward()
accum_loss += float(loss)
if (step + 1) % args.grad_accum == 0:
torch.nn.utils.clip_grad_norm_(params, 1.0)
opt.step()
opt.zero_grad()
if writer is not None:
writer.add_scalar("lm_loss", float(lm), step)
writer.add_scalar("contrastive_loss", float(contra), step)
writer.add_scalar("total_loss", float(loss), step)
if step % args.log_every == 0:
peak = (torch.cuda.max_memory_allocated() / 1e9) if torch.cuda.is_available() else 0.0
print(f" step {step:5d} lm={float(lm):.4f} contra={float(contra):.4f} "
f"total={float(loss):.4f} peak_GB={peak:.2f}")
if (step + 1) % args.ckpt_every == 0:
ckpt = os.path.join(args.ckpt_dir, f"stage1_step{step+1}.pt")
torch.save({
"encoder": model.encoder.state_dict(),
"projector": model.projector.state_dict(),
"step": step + 1,
}, ckpt)
print(f" saved {ckpt}")
step += 1
# final ckpt
ckpt = os.path.join(args.ckpt_dir, "stage1_final.pt")
torch.save({
"encoder": model.encoder.state_dict(),
"projector": model.projector.state_dict(),
"step": step,
}, ckpt)
print(f"[stage1] done. saved {ckpt} avg_loss={accum_loss/max(step,1):.4f}")
if writer is not None:
writer.close()
def main() -> None:
p = argparse.ArgumentParser(description="Stage 1 alignment training")
p.add_argument("--data", default="data/align.jsonl")
p.add_argument("--ckpt_dir", default="checkpoints/stage1")
p.add_argument("--max_steps", type=int, default=10_000)
p.add_argument("--bs", type=int, default=8)
p.add_argument("--grad_accum", type=int, default=4)
p.add_argument("--ctx", type=int, default=512)
p.add_argument("--lr", type=float, default=1e-4)
p.add_argument("--lambda_contrast", type=float, default=0.1)
p.add_argument("--perturb_std", type=float, default=0.1)
p.add_argument("--log_every", type=int, default=20)
p.add_argument("--ckpt_every", type=int, default=2000)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--tensorboard", action="store_true")
args = p.parse_args()
train(args)
if __name__ == "__main__":
main()