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:
张宗平
2026-06-30 02:11:36 +00:00
parent f509d47878
commit aec8fd5ef1
4 changed files with 333 additions and 3 deletions
+3 -2
View File
@@ -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)
+204
View File
@@ -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