From b089071bb9e61932777e94f987755c046cb4d5d1 Mon Sep 17 00:00:00 2001 From: Ange Lou Date: Wed, 30 Sep 2026 07:14:24 +0000 Subject: [PATCH 1/2] Preserve backbone precision and avoid unused inference caches Route text, media and prefix forwards through the existing adapter contract. Keep sequence hidden states in their native dtype and cast only the pointer readout vectors. Validate cache behavior and exact BF16 output/gradient parity; the CPU regression suite passes 303 tests. --- CONTRIBUTING.md | 7 +++ jevany/backbones.py | 12 ++++- jevany/model.py | 29 ++++++----- tests/test_backbone_execution.py | 86 ++++++++++++++++++++++++++++++++ 4 files changed, 121 insertions(+), 13 deletions(-) create mode 100644 tests/test_backbone_execution.py 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/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) From 95139d291c3a164ee47c997a2592535407d2d1d9 Mon Sep 17 00:00:00 2001 From: Ange Lou Date: Wed, 30 Sep 2026 07:25:29 +0000 Subject: [PATCH 2/2] Extend backbone smoke coverage for token modes and CUDA Handle existing delimiter token schemas, expose frozen-weight precision and direct-token smoke, and reject incompatible direct-mode flags. Use tiny GLM dimensions validated for BF16 SDPA training and checkpoint reload on A100. --- scripts/smoke_backbone.py | 28 ++++++++++++++++------- tests/test_backbones.py | 7 ++++-- tests/test_smoke_backbone.py | 43 ++++++++++++++++++++++++++++++++---- 3 files changed, 64 insertions(+), 14 deletions(-) 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_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])