feat(m2): MultimodalTSModel wrapper — end-to-end fwd/generate + LoRA (T2.4)
- src/tsmm/model/wrapper.py: MultimodalTSModel combining TS Encoder + Projector
+ Qwen2.5-0.5B LLM + LoRA. Two modes: freeze_llm() (stage ①) and
enable_lora(r=16) (stage ②, peft on q/v_proj). forward() takes a raw batch
{series, attributes, timestamps, events, question, answer} end-to-end
(encoder→projector→splice→LLM) OR pre-spliced inputs_embeds.
- src/tsmm/model/multimodal.py: device-aware splice (move ids/ts_embeds to
the bound embedding layer's device).
- tests/test_wrapper.py: 6 tests — finite loss, grad flows into
encoder/projector under freeze_llm, generate returns text, freeze_llm and
enable_lora mode invariants, GPU memory budget. 80 tests passing.
- Measured: stage ① fwd+bwd peak 1.87 GB (design budget ~5 GB).
Spec clarifications (small tier, comet-build Step 4):
- Wrapper batch contract = {series, attributes, timestamps, events, question,
answer} (lists of str + series tensor). Collator (T2.5) will produce this.
- Encoder/Projector kept fp32; their output cast to LLM dtype (bf16) before
splice, so inputs_embeds matches LLM weights.
This commit is contained in:
@@ -18,7 +18,7 @@
|
||||
- [x] 2.1 TS Encoder `model/ts_encoder.py`:`Patchify(patch=8, stride=4)` 通道独立切 patch 线性嵌入到 d=256;`TSEncoder(d=256, layers=2, heads=4)` 2 层 TF + 可学习位置编码;前向 `[B,T,C]`→`[B,n_patches,256]`(验证:输入 `[2,512,5]`→`[2,127,256]`;spec 澄清:n_patches=(T-P)/S+1=127 非 128;通道独立=每通道共享 patch 线性后对通道 mean-pool;参数 ≈2.1M 非 4M)
|
||||
- [x] 2.2 Projector `model/projector.py`:`Linear(256→896)+LayerNorm`(验证:输出 `[B,n_patches,896]` 对齐 Qwen2.5-0.5B hidden)
|
||||
- [x] 2.3 多模态拼接 `model/multimodal.py`:Qwen tokenizer tokenize 文本部分,TS token 作 inputs_embeds 插入 `[属性][时间戳][TS tok][事件][问题][回答][EOS]`,统一构造 inputs_embeds+attention_mask+labels(仅回答段非 -100)(验证:8 项单测过——seq_len ≤1024、mask 全 1、回答段 label 非 -100、TS token 计入序列、无 NaN、超长左截断)。spec 澄清:训练拼回答+EOS,用 `len(tokenizer)` 构建嵌入表以容纳 special tokens。
|
||||
- [ ] 2.4 训练/推理封装 `model/wrapper.py`:`MultimodalTSModel` 组合 Encoder+Projector+LLM+LoRA 挂载开关,`forward` 返 loss、`generate` 返文本,支持 `freeze_llm`(阶段①)/`enable_lora(r=16)`(阶段②)(验证:单 batch forward 3060 不 OOM;generate 能出文本)
|
||||
- [x] 2.4 训练/推理封装 `model/wrapper.py`:`MultimodalTSModel` 组合 Encoder+Projector+LLM+LoRA 挂载开关,`forward` 返 loss、`generate` 返文本,支持 `freeze_llm`(阶段①)/`enable_lora(r=16)`(阶段②)(验证:端到端 forward+backward 跑通、grad 流入 Encoder/Projector、generate 出文本、阶段①峰值 1.87GB ≪5GB 设计预算)。spec 澄清:wrapper 既接收原始 batch(series+文本列表)走端到端,也兼容预拼接 inputs_embeds;encoder/projector fp32,输出投影到 LLM dtype(bf16)。
|
||||
- [ ] 2.5 Collator `data/collator.py`:JSONL→batch 张量,处理变长 C(padding+mask)与变长文本(padding+attention_mask)(验证:batch=4 张量形状一致无 NaN)
|
||||
- [ ] 2.6 M2 出口验证:单 batch 前向+反向在 3060 跑通,loss 有限且下降趋势,阶段①模式峰值显存 ~5GB
|
||||
|
||||
|
||||
@@ -75,6 +75,7 @@ class MultimodalSplicer:
|
||||
return torch.tensor(ids, dtype=torch.long)
|
||||
|
||||
def _embed_ids(self, ids: Tensor) -> Tensor:
|
||||
ids = ids.to(self.embed.weight.device)
|
||||
return self.embed(ids) # [len, hidden]
|
||||
|
||||
# -- public API -------------------------------------------------------
|
||||
@@ -101,9 +102,9 @@ class MultimodalSplicer:
|
||||
attr_emb = self._embed_ids(attr_ids)
|
||||
ts_text_emb = self._embed_ids(ts_ids)
|
||||
if ts_embeds.shape[0] > 0:
|
||||
ts_tok_emb = ts_embeds
|
||||
ts_tok_emb = ts_embeds.to(self.embed.weight.device)
|
||||
else:
|
||||
ts_tok_emb = torch.zeros(0, self.hidden_size)
|
||||
ts_tok_emb = torch.zeros(0, self.hidden_size, device=self.embed.weight.device)
|
||||
event_emb = self._embed_ids(event_ids)
|
||||
q_emb = self._embed_ids(q_ids)
|
||||
|
||||
|
||||
@@ -0,0 +1,204 @@
|
||||
"""MultimodalTSModel (T2.4): TS Encoder + Projector + LLM + LoRA, unified wrapper.
|
||||
|
||||
Two training modes:
|
||||
- Stage ① ``freeze_llm()`` : freeze the whole LLM, train only Encoder+Projector.
|
||||
- Stage ② ``enable_lora(r)`` : keep LLM base frozen, inject LoRA on q/v_proj.
|
||||
|
||||
``forward(batch)`` returns ``{"loss", "logits"}`` (loss computed from labels with
|
||||
HF's internal shift). ``generate(batch, ...)`` returns decoded text.
|
||||
|
||||
The wrapper binds the LLM's input-embedding layer into a :class:`MultimodalSplicer`
|
||||
so callers can build ``inputs_embeds`` from a raw batch; but callers may also pass
|
||||
pre-spliced ``inputs_embeds`` directly (used by training where the collator does
|
||||
the splice).
|
||||
|
||||
Design ref: ``docs/superpowers/specs/2026-06-29-ts-as-modality-design.md`` §2, §4.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from .multimodal import MultimodalSplicer
|
||||
from .projector import Projector
|
||||
from .ts_encoder import TSEncoder
|
||||
|
||||
|
||||
class MultimodalTSModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
llm_path: str,
|
||||
encoder: TSEncoder,
|
||||
projector: Projector,
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
lora_targets: tuple = ("q_proj", "v_proj"),
|
||||
) -> None:
|
||||
super().__init__()
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
self.llm_path = llm_path
|
||||
self.dtype = dtype
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(llm_path)
|
||||
self.llm = AutoModelForCausalLM.from_pretrained(llm_path, dtype=dtype)
|
||||
self.encoder = encoder
|
||||
self.projector = projector
|
||||
self.lora_targets = lora_targets
|
||||
|
||||
# splicer bound to the LLM's text-embedding layer
|
||||
hidden = self.llm.config.hidden_size
|
||||
self.splicer = MultimodalSplicer(self.tokenizer, hidden_size=hidden)
|
||||
self.splicer.bind(self.llm.get_input_embeddings())
|
||||
|
||||
self._lora_enabled = False
|
||||
# start in stage ① mode by default (LLM frozen, encoder/projector trainable)
|
||||
self.freeze_llm()
|
||||
|
||||
# -- dtype helper -----------------------------------------------------
|
||||
@property
|
||||
def hidden_size(self) -> int:
|
||||
return self.llm.config.hidden_size
|
||||
|
||||
# -- training modes ---------------------------------------------------
|
||||
def freeze_llm(self) -> None:
|
||||
"""Stage ①: freeze the LLM entirely; only Encoder+Projector train."""
|
||||
for p in self.llm.parameters():
|
||||
p.requires_grad = False
|
||||
for p in self.encoder.parameters():
|
||||
p.requires_grad = True
|
||||
for p in self.projector.parameters():
|
||||
p.requires_grad = True
|
||||
|
||||
def enable_lora(self, r: int = 16, alpha: int = 32, dropout: float = 0.05) -> None:
|
||||
"""Stage ②: keep LLM base frozen, attach LoRA on q/v_proj.
|
||||
|
||||
Idempotent: calling twice won't double-inject adapters.
|
||||
"""
|
||||
if self._lora_enabled:
|
||||
return
|
||||
# ensure base is frozen first
|
||||
for p in self.llm.parameters():
|
||||
p.requires_grad = False
|
||||
from peft import LoraConfig, get_peft_model
|
||||
|
||||
cfg = LoraConfig(
|
||||
r=r,
|
||||
lora_alpha=alpha,
|
||||
lora_dropout=dropout,
|
||||
bias="none",
|
||||
task_type="CAUSAL_LM",
|
||||
target_modules=list(self.lora_targets),
|
||||
)
|
||||
self.llm = get_peft_model(self.llm, cfg)
|
||||
# re-bind splicer to the (possibly wrapped) embedding layer
|
||||
self.splicer.bind(self.llm.get_input_embeddings())
|
||||
self._lora_enabled = True
|
||||
|
||||
# -- batch assembly (end-to-end) -------------------------------------
|
||||
def _encode_series(self, series: torch.Tensor) -> torch.Tensor:
|
||||
"""series [B, T, C] → projected TS tokens [B, n_patches, hidden]."""
|
||||
ts = self.encoder(series) # [B, n_patches, d]
|
||||
ts = self.projector(ts) # [B, n_patches, hidden]
|
||||
return ts.to(self.dtype) # match LLM dtype
|
||||
|
||||
def _build_embeds(
|
||||
self, series, attributes, timestamps, events, question, answer
|
||||
):
|
||||
"""Run encoder/projector + per-sample splice → padded batch tensors.
|
||||
|
||||
Returns (inputs_embeds, attention_mask, labels) on the LLM device.
|
||||
"""
|
||||
device = self.llm.device if hasattr(self.llm, "device") else next(self.llm.parameters()).device
|
||||
series = series.to(device)
|
||||
ts_tokens = self._encode_series(series) # [B, n_patches, hidden]
|
||||
B = series.shape[0]
|
||||
ans_provided = answer is not None
|
||||
samples = []
|
||||
for i in range(B):
|
||||
out = self.splicer.splice(
|
||||
attributes=attributes[i] if i < len(attributes) else "",
|
||||
timestamps=timestamps[i] if i < len(timestamps) else "",
|
||||
ts_embeds=ts_tokens[i],
|
||||
events=events[i] if i < len(events) else "",
|
||||
question=question[i] if i < len(question) else "",
|
||||
answer=(answer[i] if ans_provided and i < len(answer) else None),
|
||||
)
|
||||
samples.append((out.inputs_embeds[0], out.attention_mask[0], out.labels[0]))
|
||||
max_seq = max(s[0].shape[0] for s in samples)
|
||||
hidden = samples[0][0].shape[1]
|
||||
ie = torch.zeros(B, max_seq, hidden, device=device, dtype=samples[0][0].dtype)
|
||||
attn = torch.zeros(B, max_seq, dtype=torch.long, device=device)
|
||||
labels = torch.full((B, max_seq), -100, dtype=torch.long, device=device)
|
||||
for i, (e, a, l) in enumerate(samples):
|
||||
L = e.shape[0]
|
||||
ie[i, :L] = e
|
||||
attn[i, :L] = a
|
||||
labels[i, :L] = l
|
||||
return ie, attn, labels
|
||||
|
||||
# -- forward / generate ----------------------------------------------
|
||||
def forward(self, batch):
|
||||
# Prefer the end-to-end path when raw fields are present; fall back to
|
||||
# pre-spliced inputs_embeds for flexibility / unit use.
|
||||
if "series" in batch:
|
||||
ie, attn, labels = self._build_embeds(
|
||||
batch["series"],
|
||||
batch.get("attributes", []),
|
||||
batch.get("timestamps", []),
|
||||
batch.get("events", []),
|
||||
batch.get("question", []),
|
||||
batch.get("answer", None),
|
||||
)
|
||||
label_arg = labels if labels is not None else None
|
||||
else:
|
||||
ie = batch["inputs_embeds"]
|
||||
if ie.dtype != self.dtype:
|
||||
ie = ie.to(self.dtype)
|
||||
attn = batch["attention_mask"]
|
||||
label_arg = batch.get("labels")
|
||||
out = self.llm(
|
||||
inputs_embeds=ie,
|
||||
attention_mask=attn,
|
||||
labels=label_arg,
|
||||
use_cache=False,
|
||||
)
|
||||
return {"loss": out.loss, "logits": out.logits}
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(
|
||||
self,
|
||||
batch,
|
||||
max_new_tokens: int = 64,
|
||||
do_sample: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> List[str]:
|
||||
if "series" in batch:
|
||||
ie, attn, _ = self._build_embeds(
|
||||
batch["series"],
|
||||
batch.get("attributes", []),
|
||||
batch.get("timestamps", []),
|
||||
batch.get("events", []),
|
||||
batch.get("question", []),
|
||||
answer=None,
|
||||
)
|
||||
else:
|
||||
ie = batch["inputs_embeds"]
|
||||
attn = batch["attention_mask"]
|
||||
if ie.dtype != self.dtype:
|
||||
ie = ie.to(self.dtype)
|
||||
gen_ids = self.llm.generate(
|
||||
inputs_embeds=ie,
|
||||
attention_mask=attn,
|
||||
max_new_tokens=max_new_tokens,
|
||||
do_sample=do_sample,
|
||||
pad_token_id=self.tokenizer.eos_token_id,
|
||||
eos_token_id=self.tokenizer.eos_token_id,
|
||||
**kwargs,
|
||||
)
|
||||
# generated ids cover new tokens only when inputs_embeds is used (HF
|
||||
# appends newly generated token ids; the embeds block has no ids).
|
||||
# Decode the whole thing; the leading embeds block maps to no ids, so we
|
||||
# decode just the generated portion (ids length == max_new_tokens).
|
||||
texts = self.tokenizer.batch_decode(gen_ids, skip_special_tokens=True)
|
||||
return texts
|
||||
@@ -0,0 +1,125 @@
|
||||
"""Tests for MultimodalTSModel wrapper (T2.4).
|
||||
|
||||
Verifies the end-to-end path: series → Encoder → Projector → splice → LLM.
|
||||
- forward returns a finite scalar loss
|
||||
- generate returns text
|
||||
- freeze_llm=True (stage ①): LLM frozen, encoder+projector trainable, and the
|
||||
loss carries grad into encoder/projector (backward works)
|
||||
- enable_lora(r=16) (stage ②): base LLM weights frozen, LoRA adapters trainable
|
||||
|
||||
GPU (RTX 3060) forward/backward + peak-memory checks live here too.
|
||||
"""
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from tsmm.model.ts_encoder import TSEncoder
|
||||
from tsmm.model.projector import Projector
|
||||
from tsmm.model.wrapper import MultimodalTSModel
|
||||
|
||||
LLM_PATH = "/home/zhangzp/models/Qwen2.5-0.5B-Instruct"
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(), reason="needs CUDA for LLM forward"
|
||||
)
|
||||
|
||||
|
||||
def make_model():
|
||||
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=torch.bfloat16,
|
||||
)
|
||||
|
||||
|
||||
def make_raw_batch(batch_size=2, T=512, C=5):
|
||||
"""Raw batch contract (what the collator T2.5 will produce)."""
|
||||
series = torch.randn(batch_size, T, C)
|
||||
return {
|
||||
"series": series,
|
||||
"attributes": ["CPU 内存"] * batch_size,
|
||||
"timestamps": ["t0..t1"] * batch_size,
|
||||
"events": [""] * batch_size,
|
||||
"question": ["是否存在异常?"] * batch_size,
|
||||
"answer": ["正常。"] * batch_size,
|
||||
}
|
||||
|
||||
|
||||
class TestForwardGenerate:
|
||||
def test_forward_returns_finite_loss(self):
|
||||
model = make_model().cuda()
|
||||
batch = make_raw_batch(batch_size=2)
|
||||
with torch.no_grad():
|
||||
out = model(batch)
|
||||
assert "loss" in out
|
||||
assert torch.isfinite(out["loss"])
|
||||
assert out["loss"].dim() == 0
|
||||
|
||||
def test_forward_backward_flows_into_encoder_projector(self):
|
||||
# stage ①: LLM frozen, so the only grad path is encoder+projector
|
||||
model = make_model().cuda()
|
||||
model.freeze_llm()
|
||||
model.train()
|
||||
batch = make_raw_batch(batch_size=1)
|
||||
out = model(batch)
|
||||
out["loss"].backward()
|
||||
# encoder/projector must have received gradients
|
||||
for p in model.encoder.parameters():
|
||||
assert p.grad is not None
|
||||
for p in model.projector.parameters():
|
||||
assert p.grad is not None
|
||||
|
||||
def test_generate_returns_text(self):
|
||||
model = make_model().cuda()
|
||||
batch = make_raw_batch(batch_size=1)
|
||||
text = model.generate(batch, max_new_tokens=8)
|
||||
assert isinstance(text, list)
|
||||
assert len(text) == 1
|
||||
assert isinstance(text[0], str)
|
||||
|
||||
|
||||
class TestStageModes:
|
||||
def test_freeze_llm_only_trains_encoder_projector(self):
|
||||
model = make_model()
|
||||
model.freeze_llm() # stage ①
|
||||
llm_requires_grad = [p.requires_grad for p in model.llm.parameters()]
|
||||
assert all(not r for r in llm_requires_grad)
|
||||
enc_proj_requires_grad = (
|
||||
[p.requires_grad for p in model.encoder.parameters()]
|
||||
+ [p.requires_grad for p in model.projector.parameters()]
|
||||
)
|
||||
assert all(enc_proj_requires_grad)
|
||||
|
||||
def test_enable_lora_freezes_base_injects_adapters(self):
|
||||
model = make_model()
|
||||
model.freeze_llm()
|
||||
model.enable_lora(r=16) # stage ②
|
||||
# base (non-LoRA) LLM weights must stay frozen
|
||||
base_frozen = all(
|
||||
not p.requires_grad
|
||||
for n, p in model.llm.named_parameters()
|
||||
if "lora_" not in n
|
||||
)
|
||||
assert base_frozen
|
||||
# LoRA adapter params must be trainable
|
||||
lora_trainable = [
|
||||
p for n, p in model.llm.named_parameters()
|
||||
if "lora_" in n and p.requires_grad
|
||||
]
|
||||
assert len(lora_trainable) > 0
|
||||
# there must be lora-named modules
|
||||
lora_names = [n for n, _ in model.named_modules() if "lora" in n.lower()]
|
||||
assert len(lora_names) > 0
|
||||
|
||||
|
||||
class TestGPUMemory:
|
||||
def test_forward_backward_no_oom_and_reasonable_mem(self):
|
||||
torch.cuda.reset_peak_memory_stats()
|
||||
model = make_model().cuda()
|
||||
model.freeze_llm() # stage ① mode for memory check
|
||||
model.train()
|
||||
batch = make_raw_batch(batch_size=2)
|
||||
out = model(batch)
|
||||
out["loss"].backward()
|
||||
peak_gb = torch.cuda.max_memory_allocated() / 1e9
|
||||
# stage ① should be well under the design M2 target (~5GB); use 8GB headroom
|
||||
assert peak_gb < 8.0, f"peak {peak_gb:.2f} GB"
|
||||
Reference in New Issue
Block a user