diff --git a/.gitignore b/.gitignore index edd2d9b..08a59a5 100644 --- a/.gitignore +++ b/.gitignore @@ -126,3 +126,8 @@ coverage/ .context-*/ tmp/ .temp/ + +# local artifacts (design §6.1) +/checkpoints/ +/data/*.jsonl + diff --git a/openspec/changes/ts-as-modality/tasks.md b/openspec/changes/ts-as-modality/tasks.md index 0c46463..c4d8f3a 100644 --- a/openspec/changes/ts-as-modality/tasks.md +++ b/openspec/changes/ts-as-modality/tasks.md @@ -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+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.4 M3 出口验证:阶段① ckpt 产出;扰动检验通过;峰值显存 ≤6GB(失败回查对比损失权重/数据质量) diff --git a/scripts/train_stage1.sh b/scripts/train_stage1.sh new file mode 100755 index 0000000..a8e6567 --- /dev/null +++ b/scripts/train_stage1.sh @@ -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 diff --git a/src/tsmm/data/collator.py b/src/tsmm/data/collator.py index 0dff93f..56355ae 100644 --- a/src/tsmm/data/collator.py +++ b/src/tsmm/data/collator.py @@ -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: diff --git a/src/tsmm/train/stage1.py b/src/tsmm/train/stage1.py new file mode 100644 index 0000000..7703abb --- /dev/null +++ b/src/tsmm/train/stage1.py @@ -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()