Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions CONTRIBUTING.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 |
Expand Down
12 changes: 11 additions & 1 deletion jevany/backbones.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
29 changes: 17 additions & 12 deletions jevany/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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)
Expand All @@ -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))
Expand All @@ -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):
Expand All @@ -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)]

Expand All @@ -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):
Expand All @@ -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())

Expand All @@ -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)]
Expand Down
28 changes: 20 additions & 8 deletions scripts/smoke_backbone.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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"))
Expand All @@ -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),
Expand Down Expand Up @@ -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,
Expand All @@ -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)
Expand Down
86 changes: 86 additions & 0 deletions tests/test_backbone_execution.py
Original file line number Diff line number Diff line change
@@ -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)
7 changes: 5 additions & 2 deletions tests/test_backbones.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"],
Expand Down
Loading
Loading