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:
@@ -126,3 +126,8 @@ coverage/
|
|||||||
.context-*/
|
.context-*/
|
||||||
tmp/
|
tmp/
|
||||||
.temp/
|
.temp/
|
||||||
|
|
||||||
|
# local artifacts (design §6.1)
|
||||||
|
/checkpoints/
|
||||||
|
/data/*.jsonl
|
||||||
|
|
||||||
|
|||||||
@@ -25,7 +25,7 @@
|
|||||||
## 3. M3 · 阶段①对齐训练
|
## 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 强扰动检验另行评估。
|
- [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+Projector(AdamW, 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+Projector(AdamW, lr=1e-4),bs=8/grad_accum=4/ctx=512/BF16/梯度检查点,每 2000 step 存 ckpt + TensorBoard,loss=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.3 TS-token 有效性检验:held-out 原时序 vs 扰动时序回答变化率 + 「必看时序」样本答对率对照纯 LLM(验证:扰动后变化率 >50%,必看时序答对率显著高于纯 LLM)
|
||||||
- [ ] 3.4 M3 出口验证:阶段① ckpt 产出;扰动检验通过;峰值显存 ≤6GB(失败回查对比损失权重/数据质量)
|
- [ ] 3.4 M3 出口验证:阶段① ckpt 产出;扰动检验通过;峰值显存 ≤6GB(失败回查对比损失权重/数据质量)
|
||||||
|
|
||||||
|
|||||||
Executable
+23
@@ -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
|
||||||
@@ -64,7 +64,11 @@ class Collator:
|
|||||||
def _stack_series(self, samples: List[dict]) -> torch.Tensor:
|
def _stack_series(self, samples: List[dict]) -> torch.Tensor:
|
||||||
tensors = []
|
tensors = []
|
||||||
for s in samples:
|
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
|
T, C = arr.shape
|
||||||
# truncate / pad T to max_T
|
# truncate / pad T to max_T
|
||||||
if T >= self.max_T:
|
if T >= self.max_T:
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user