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