diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index b8b4e9e..0cde838 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -42,6 +42,13 @@ Keep JSON request/response compatibility when changing inference. Add labels onl to training records and keep them out of model-facing inputs. Data converters should record source revisions and preserve evaluation separation. +For another model family, extend `BackboneAdapter` and select it with +`--backbone-adapter module:Class`. Its `forward(model, **inputs)` hook serves text, +native media and prefix reuse; return `last_hidden_state` in the native dtype and +`past_key_values` when caching is explicitly requested. Ordinary scoring disables +caching. Declare packed-mask and prefix-cache support separately, and validate +training, checkpoint reload and serving before enabling either capability. + The code is organized around user entry points: | Location | Responsibility | diff --git a/jevany/backbones.py b/jevany/backbones.py index c08ff14..09d9627 100644 --- a/jevany/backbones.py +++ b/jevany/backbones.py @@ -287,13 +287,23 @@ def new_cache(self, model: PreTrainedModel) -> DynamicCache: """Create a cache supporting deepcopy/reorder and, for packed mode, crop.""" return DynamicCache(config=model.config) + def forward(self, model, **inputs): + """Return native hidden states and, only when requested, a reusable cache. + + Override here to adapt a family's forward signature or output layout. + Ordinary decision scoring does not need a generation cache; prefix + methods explicitly request one with use_cache=True. + """ + inputs.setdefault("use_cache", False) + return model(**inputs) + def encode_media(self, processor, record, **kwargs) -> dict: raise ValueError("this backbone adapter does not support media") def forward_media(self, language_model, multimodal_model, inputs: dict) -> torch.Tensor: """Return token hidden states, retaining each processor's model inputs.""" model = multimodal_model if multimodal_model is not None else language_model - return model(**inputs).last_hidden_state + return self.forward(model, **inputs).last_hidden_state def _video_inputs(processor, paths: list[str], *, num_frames: int) -> dict: diff --git a/jevany/model.py b/jevany/model.py index a2bcde9..0001d74 100644 --- a/jevany/model.py +++ b/jevany/model.py @@ -467,14 +467,14 @@ def _pad_rows(self, rows): return ids, pos, att def hidden_batch(self, encs): - """[B, L_max, d] hidden states for a right-padded batch of encoded records under the packed block-causal mask.""" + """[B, L_max, d] packed hidden states in the backbone's native dtype.""" ids, pos, _ = self._pad_rows([(e["ids"], e["pos"]) for e in encs]) isolate = any(e.get("option_isolation") for e in encs) if isolate and not all(e.get("option_isolation") for e in encs): raise ValueError("cannot mix option-isolated and plain encodings in one batch") lm_dtype = next(self.lm.parameters()).dtype mask = branch_mask_batch([e["seg"] for e in encs], self.device, dtype=lm_dtype, opts=[e["opt"] for e in encs] if isolate else None, length=ids.shape[1]) - return self.lm(input_ids=ids, position_ids=pos, attention_mask=mask).last_hidden_state.float() # head stays fp32 + return self.adapter.forward(self.lm, input_ids=ids, position_ids=pos, attention_mask=mask).last_hidden_state def _question_readout(self, h, decide, options): query = h[decide] @@ -485,7 +485,8 @@ def _question_readout(self, h, decide, options): candidates = torch.tensor(self.verbalizer_ids[:len(options)], device=logits.device) logits = logits.index_select(0, candidates) return logits if self.temperature == 1.0 else logits / self.temperature - return self.head(query, h[torch.tensor(options, device=self.device)]) + # The pointer head stays FP32; only its selected vectors need upcasting. + return self.head(query.float(), h[torch.tensor(options, device=self.device)].float()) def _readout(self, h, enc): return [self._question_readout(h, d, oi) @@ -506,7 +507,7 @@ def forward_rows_batch(self, encs): for r in brs: rows.append((S + r["ids"], Sp + r["pos"])); readouts.append((b, len(S) + r["decide"], [len(S) + o for o in r["opts"]])) ids, pos, att = self._pad_rows(rows) - h = self.lm(input_ids=ids, position_ids=pos, attention_mask=att).last_hidden_state.float() + h = self.adapter.forward(self.lm, input_ids=ids, position_ids=pos, attention_mask=att).last_hidden_state out = [[] for _ in encs] for i, (b, d, oi) in enumerate(readouts): out[b].append(self._question_readout(h[i], d, oi)) @@ -522,7 +523,7 @@ def forward_multimodal(self, enc): kwargs = media_to(enc["mm"], self.device) ids = torch.tensor([enc["ids"]], device=self.device) kwargs.setdefault("attention_mask", torch.ones_like(ids)) - hidden = self.adapter.forward_media(self.lm, self.mm, {"input_ids": ids, **kwargs})[0].float() + hidden = self.adapter.forward_media(self.lm, self.mm, {"input_ids": ids, **kwargs})[0] return self._readout(hidden, enc) def forward_batch(self, encs): @@ -548,7 +549,8 @@ def _branch_rows_from_prefix(self, enc, cache): cache = copy.deepcopy(cache); cache.reorder_cache(torch.zeros(Q, dtype=torch.long, device=self.device)) ids, pos, att = self._pad_rows([(r["ids"], r["pos"]) for r in rows]) att = torch.cat([torch.ones((Q, len(S)), dtype=torch.long, device=self.device), att], 1) # the cached state tokens are all real - h = self.lm(input_ids=ids, position_ids=pos, attention_mask=att, past_key_values=cache, use_cache=True).last_hidden_state.float() + h = self.adapter.forward(self.lm, input_ids=ids, position_ids=pos, attention_mask=att, + past_key_values=cache, use_cache=True).last_hidden_state return [F.softmax(self._question_readout(h[i], r["decide"], r["opts"]), -1).cpu() for i, r in enumerate(rows)] @@ -565,8 +567,9 @@ def prefix(self, enc): Ls = enc["seg"].count(0) ids = torch.tensor([enc["ids"][:Ls]], device=self.device); pos = torch.tensor([enc["pos"][:Ls]], device=self.device) # the cache must know the layer types (hybrid backbones keep recurrent + conv states per DeltaNet layer) - out = self.lm(input_ids=ids, position_ids=pos, past_key_values=self.adapter.new_cache(self.lm), use_cache=True) - return Ls, out.past_key_values, out.last_hidden_state[0].float() + out = self.adapter.forward(self.lm, input_ids=ids, position_ids=pos, + past_key_values=self.adapter.new_cache(self.lm), use_cache=True) + return Ls, out.past_key_values, out.last_hidden_state[0] @torch.no_grad() def probs_and_prefix(self, enc): @@ -582,8 +585,9 @@ def probs_and_prefix(self, enc): ids = torch.tensor([enc["ids"]], device=self.device); pos = torch.tensor([enc["pos"]], device=self.device) dt = next(self.lm.parameters()).dtype mask = branch_mask_batch([enc["seg"]], self.device, dtype=dt, opts=[enc["opt"]] if enc.get("option_isolation") else None) - out = self.lm(input_ids=ids, position_ids=pos, attention_mask=mask, past_key_values=self.adapter.new_cache(self.lm), use_cache=True) - h = out.last_hidden_state[0].float() + out = self.adapter.forward(self.lm, input_ids=ids, position_ids=pos, attention_mask=mask, + past_key_values=self.adapter.new_cache(self.lm), use_cache=True) + h = out.last_hidden_state[0] out.past_key_values.crop(-(len(enc["ids"]) - Ls)) # keep the state only (negative = drop that many trailing tokens; positive form deprecated in transformers 5) return [F.softmax(z, -1).cpu() for z in self._readout(h, enc)], (Ls, out.past_key_values, h[:Ls].clone()) @@ -600,8 +604,9 @@ def probs_with_prefix(self, enc, prefix): dt = next(self.lm.parameters()).dtype mask = branch_mask_batch([enc["seg"]], self.device, dtype=dt, opts=[enc["opt"]] if enc.get("option_isolation") else None)[:, :, Ls:, :] try: - out = self.lm(input_ids=ids, position_ids=pos, past_key_values=cache, attention_mask=mask, use_cache=True) - h = torch.cat([h_state, out.last_hidden_state[0].float()], 0) + out = self.adapter.forward(self.lm, input_ids=ids, position_ids=pos, + past_key_values=cache, attention_mask=mask, use_cache=True) + h = torch.cat([h_state, out.last_hidden_state[0]], 0) finally: cache.crop(-(len(enc["ids"]) - Ls)) return [F.softmax(z, -1).cpu() for z in self._readout(h, enc)] diff --git a/scripts/smoke_backbone.py b/scripts/smoke_backbone.py index 41106e0..7e352b8 100644 --- a/scripts/smoke_backbone.py +++ b/scripts/smoke_backbone.py @@ -151,10 +151,14 @@ def main(argv: list[str] | None = None) -> None: parser.add_argument("--out", type=Path, required=True) parser.add_argument("--steps", type=int, default=12) parser.add_argument("--lr", type=float, default=0.0002) - parser.add_argument("--head-lr", type=float, default=0.0001) + parser.add_argument("--head-lr", type=float, + help="pointer head learning rate (default: 0.0001); lm_token requires 0") parser.add_argument("--device", choices=["cpu", "cuda"], default="cuda") parser.add_argument("--dtype", choices=["fp32", "bf16"], help="training autocast precision; defaults to bf16 on CUDA and fp32 on CPU") + parser.add_argument("--weights-dtype", choices=["fp32", "bf16"], + help="frozen backbone precision; defaults to bf16 on CUDA and fp32 on CPU") + parser.add_argument("--decision-mode", choices=["pointer", "lm_token"], default="pointer") parser.add_argument("--branch-mode", choices=["auto", "packed", "rows"], default="auto") parser.add_argument("--media", choices=["image", "video"], help="exercise native media training") parser.add_argument("--mixed-text", action="store_true", help="include text-only rows in a media fixture") @@ -166,8 +170,12 @@ def main(argv: list[str] | None = None) -> None: if args.rlcr and args.init_from is None: parser.error("--rlcr requires --init-from") training_dtype = args.dtype or ("bf16" if args.device == "cuda" else "fp32") + weights_dtype = args.weights_dtype or ("bf16" if args.device == "cuda" else "fp32") + head_lr = args.head_lr if args.head_lr is not None else (0 if args.decision_mode == "lm_token" else 0.0001) if training_dtype == "bf16" and args.device != "cuda": parser.error("--dtype bf16 requires --device cuda") + if args.decision_mode == "lm_token" and (args.rlcr or head_lr): + parser.error("lm_token requires supervised training and --head-lr 0") rank = int(os.environ.get("RANK", "0")) world_size = int(os.environ.get("WORLD_SIZE", "1")) local_rank = int(os.environ.get("LOCAL_RANK", "0")) @@ -193,10 +201,10 @@ def main(argv: list[str] | None = None) -> None: command = [ "--base", args.base, "--base-revision", args.revision, "--data", str(fixture / "train.jsonl"), "--out", str(checkpoint_path), - "--device", args.device, "--weights-dtype", "bf16" if args.device == "cuda" else "fp32", - "--dtype", training_dtype, + "--device", args.device, "--weights-dtype", weights_dtype, + "--dtype", training_dtype, "--decision-mode", args.decision_mode, "--epochs", str(args.steps), "--max-steps", str(args.steps), "--batch", "2", "--accum", "1", - "--lora", "4", "--head-dim", "32", "--lr", str(args.lr), "--head-lr", str(args.head_lr), + "--lora", "4", "--head-dim", "32", "--lr", str(args.lr), "--head-lr", str(head_lr), "--checkpointing", "1", "--weight-decay", "0", "--seed", "17", "--p-none", "0", "--p-none-distract", "0", "--p-distract", "0", "--branch-mode", args.branch_mode, "--eval-suite", str(fixture), @@ -270,14 +278,18 @@ def main(argv: list[str] | None = None) -> None: lora_updated = any("lora_" in name and not torch.equal(value, initial[name]) for name, value in weights.items()) serving = check_serving(local, fixture) + token_schema = tok.init_kwargs.get( + "jevany_token_schema", "explicit" if "jevany_decision_tokens" in tok.init_kwargs else "legacy", + ) report = { "base": args.base, "base_revision": args.revision, "model_type": model.lm.config.model_type, - "branch_mode": model.branch_mode, "token_schema": tok.init_kwargs["jevany_token_schema"], + "branch_mode": model.branch_mode, "token_schema": token_schema, "decision_token_ids": [tok.convert_tokens_to_ids(token) for token in decision_tokens(tok)], "backbone_adapter": model.backbone_adapter, "media": args.media, "mixed_text": args.mixed_text, "trained_delimiter_embeddings": model.special_embeddings, "steps": args.steps, - "learning_rate": args.lr, "head_learning_rate": args.head_lr, - "training_dtype": training_dtype, + "learning_rate": args.lr, "head_learning_rate": head_lr, + "training_dtype": training_dtype, "weights_dtype": weights_dtype, + "decision_mode": model.decision_mode, "world_size": world_size, "losses": losses, "adapter_tensors_finite": finite, "evaluation_records": len(records), "lora_updated": lora_updated, "checkpoint_probability_max_delta": maximum_delta, @@ -286,7 +298,7 @@ def main(argv: list[str] | None = None) -> None: "prefix_cache_supported": prefix_supported, "media_probability_max_delta": media_delta, "training": read_json(checkpoint_path / "training_metrics.json"), - "scope": "Real pretrained weights; repeated tiny training fixture; no generalization claim.", + "scope": "Repeated tiny training fixture for optimization and serialization; no generalization claim.", } report["passed"] = ( finite and lora_updated and all(math.isfinite(point["nll"]) for point in losses) diff --git a/tests/test_backbone_execution.py b/tests/test_backbone_execution.py new file mode 100644 index 0000000..c235e31 --- /dev/null +++ b/tests/test_backbone_execution.py @@ -0,0 +1,86 @@ +"""Native precision, gradient parity and the adapter's forward/cache contract.""" +import pytest +import torch + +from jevany.backbones import BackboneAdapter +from jevany.model import DecisionModel, load_tokenizer +from test_backbones import RECORD, make_base + + +@pytest.fixture(autouse=True) +def single_thread(): + previous = torch.get_num_threads() + torch.set_num_threads(1) + yield + torch.set_num_threads(previous) + + +class RecordingAdapter(BackboneAdapter): + def __init__(self): + self.calls = [] + + def forward(self, model, **inputs): + output = super().forward(model, **inputs) + self.calls.append((bool(inputs.get("use_cache")), output.past_key_values is not None, + output.last_hidden_state.dtype)) + return output + + +@pytest.mark.parametrize("branch_mode", ["packed", "rows"]) +def test_scoring_and_prefix_reuse_go_through_adapter(tmp_path, branch_mode): + base = tmp_path / "base" + make_base(base, "llama") + tokenizer = load_tokenizer(base) + model = DecisionModel(base, tokenizer, "cpu", lora=2, head_dim=8, branch_mode=branch_mode) + model.adapter = RecordingAdapter() + model.eval() + encoded = model.encode(tokenizer, RECORD) + # The base enables caching by default; ordinary decision scoring opts out. + assert model.lm.config.use_cache + expected = model.probs(encoded) + assert model.adapter.calls == [(False, False, torch.float32)] + first, prefix = model.probs_and_prefix(encoded) + for _ in range(2): + reused = model.probs_with_prefix(encoded, prefix) + for direct, miss, hit in zip(expected, first, reused): + torch.testing.assert_close(direct, miss, atol=2e-6, rtol=1e-5) + torch.testing.assert_close(direct, hit, atol=2e-6, rtol=1e-5) + assert all(requested and returned for requested, returned, _ in model.adapter.calls[1:]) + + +@pytest.mark.parametrize("branch_mode", ["packed", "rows"]) +def test_bf16_readout_and_gradients_match_full_hidden_upcast(tmp_path, branch_mode): + torch.manual_seed(17) + base = tmp_path / "base" + make_base(base, "llama") + tokenizer = load_tokenizer(base) + model = DecisionModel(base, tokenizer, "cpu", lora=2, head_dim=8, lora_dropout=0, + dtype=torch.bfloat16, branch_mode=branch_mode) + model.adapter = RecordingAdapter() + model.train() + encoded = model.encode(tokenizer, RECORD) + logits = model(encoded) + sum(value.square().sum() for value in logits).backward() + gradients = {name: parameter.grad.clone() for name, parameter in model.named_parameters() + if parameter.grad is not None} + assert gradients and all(torch.isfinite(value).all() for value in gradients.values()) + assert model.adapter.calls == [(False, False, torch.bfloat16)] + if branch_mode == "packed": + with torch.no_grad(): + assert model.hidden_batch([encoded]).dtype == torch.bfloat16 + model.zero_grad(set_to_none=True) + + # Emulate the old full-sequence FP32 conversion using the same frozen base, + # adapter weights and head, then compare both outputs and backpropagation. + def full_upcast(module, args, output): + output.last_hidden_state = output.last_hidden_state.float() + return output + + with model.lm.register_forward_hook(full_upcast): + reference = model(encoded) + sum(value.square().sum() for value in reference).backward() + for actual, expected in zip(logits, reference): + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + for name, parameter in model.named_parameters(): + if name in gradients: + torch.testing.assert_close(parameter.grad, gradients[name], rtol=0, atol=0) diff --git a/tests/test_backbones.py b/tests/test_backbones.py index 5e25b28..3558e06 100644 --- a/tests/test_backbones.py +++ b/tests/test_backbones.py @@ -56,10 +56,13 @@ def make_base(path, family, *, legacy=False): rope_parameters={k: {"rope_type": "default", "rope_theta": 10000} for k in ("sliding_attention", "full_attention")}), "mistral": lambda: MistralConfig(**common, sliding_window=16), + # The 4/4/8 head fixture faults in torch 2.8 BF16 SDPA on A100. + # These dimensions also pass CUDA training and checkpoint reload. "glm": lambda: Glm4MoeLiteConfig( - **{**common, "num_key_value_heads": 4}, moe_intermediate_size=16, + **{**common, "hidden_size": 64, "intermediate_size": 128, "num_key_value_heads": 4}, + moe_intermediate_size=32, n_routed_experts=4, num_experts_per_tok=2, q_lora_rank=16, - kv_lora_rank=8, qk_nope_head_dim=4, qk_rope_head_dim=4, v_head_dim=8, + kv_lora_rank=16, qk_nope_head_dim=16, qk_rope_head_dim=16, v_head_dim=16, ), "nemotron": lambda: NemotronHConfig( **common, head_dim=8, layers_block_type=["mamba", "attention", "moe"], diff --git a/tests/test_smoke_backbone.py b/tests/test_smoke_backbone.py index 9e88657..5000b14 100644 --- a/tests/test_smoke_backbone.py +++ b/tests/test_smoke_backbone.py @@ -27,15 +27,15 @@ def test_mixed_fixture_evaluates_media_and_text_rows(tmp_path): assert original.get("media") == reversed_row.get("media") -@pytest.mark.parametrize("family", ["llama", "gemma4_unified"]) -def test_sft_rlcr_and_serving_smoke(tmp_path, family): +@pytest.mark.parametrize("family,legacy", [("llama", False), ("llama", True), ("gemma4_unified", False)]) +def test_sft_rlcr_and_serving_smoke(tmp_path, family, legacy): previous = torch.get_num_threads() torch.set_num_threads(1) try: torch.manual_seed(17) base = tmp_path / "base" if family == "llama": - make_base(base, family) + make_base(base, family, legacy=legacy) else: make_vision_base(base, family) args = ["--base", str(base), "--device", "cpu", "--lr", "0.001", "--head-lr", "0.001"] @@ -51,6 +51,10 @@ def test_sft_rlcr_and_serving_smoke(tmp_path, family): assert report["passed"] assert report["objective"] == stage assert report["training_dtype"] == "fp32" + assert report["weights_dtype"] == "fp32" + assert report["decision_mode"] == "pointer" + if legacy: + assert report["token_schema"] == "explicit" assert report["lora_updated"] assert report["serving"]["passed"] assert report["serving"]["invalid_requests_rejected"] @@ -58,7 +62,38 @@ def test_sft_rlcr_and_serving_smoke(tmp_path, family): torch.set_num_threads(previous) -@pytest.mark.parametrize("extra", [["--steps", "0"], ["--rlcr"], ["--device", "cpu", "--dtype", "bf16"]]) +def test_direct_token_train_reload_and_serving_smoke(tmp_path): + from transformers import AutoModelForCausalLM, AutoTokenizer + + previous = torch.get_num_threads() + torch.set_num_threads(1) + try: + torch.manual_seed(17) + base = tmp_path / "base" + make_base(base, "llama") + tokenizer = AutoTokenizer.from_pretrained(base) + # The public direct-token contract needs 255 single-token verbalizers. + tokenizer.add_tokens([chr(index) for index in range(33, 512) if chr(index).isprintable()]) + model = AutoModelForCausalLM.from_pretrained(base) + model.resize_token_embeddings(len(tokenizer), mean_resizing=False) + model.save_pretrained(base) + tokenizer.save_pretrained(base) + main(["--base", str(base), "--device", "cpu", "--decision-mode", "lm_token", + "--weights-dtype", "fp32", "--steps", "12", "--lr", "0.001", + "--out", str(tmp_path / "direct")]) + report = json.loads((tmp_path / "direct/smoke.json").read_text()) + assert report["passed"] + assert report["decision_mode"] == "lm_token" + assert report["head_learning_rate"] == 0 + assert report["serving"]["python_http_answers_equal"] + finally: + torch.set_num_threads(previous) + + +@pytest.mark.parametrize("extra", [ + ["--steps", "0"], ["--rlcr"], ["--device", "cpu", "--dtype", "bf16"], + ["--decision-mode", "lm_token", "--head-lr", "0.001"], +]) def test_invalid_smoke_arguments(tmp_path, extra): with pytest.raises(SystemExit): main(["--base", "unused", "--out", str(tmp_path), *extra])