"""Native Qwen3.5 vision-language backbone with candidate scalar readout.""" from pathlib import Path import torch from peft import LoraConfig, PeftModel, get_peft_model, prepare_model_for_kbit_training from safetensors.torch import load_file, save_file from torch import nn from transformers import AutoProcessor, BitsAndBytesConfig, Qwen3_5ForConditionalGeneration from decision_schema import options from input_contract import normalize from training_schema import validate_question from storage import load, write from templates import messages class GeneralDecisionModel(nn.Module): def __init__(self, base, *, mode="lora", rank=16, adapter=None, device="cuda", quantization=None, candidate_batch_size=1, delta_backend="torch"): super().__init__() self.base_path = str(Path(base).resolve()) self.input_format = "legacy" if adapter and (Path(adapter) / "config.json").exists(): saved = load(Path(adapter) / "config.json") self.input_format = saved.get("input_format", "legacy") quantization = saved.get("quantization", quantization) candidate_batch_size = saved.get("candidate_batch_size", candidate_batch_size) delta_backend = saved.get("delta_backend", delta_backend) self.quantization = quantization self.candidate_batch_size = candidate_batch_size self.delta_backend = delta_backend if delta_backend == "fla": from fla.ops.gated_delta_rule import chunk_gated_delta_rule from transformers.models.qwen3_5 import modeling_qwen3_5 # Same mathematical operation and Q/K normalization flag. Pinned FLA # kernels support backward on GB10; avoid the slow torch chunk loop. modeling_qwen3_5.torch_chunk_gated_delta_rule = chunk_gated_delta_rule elif delta_backend != "torch": raise ValueError("unknown DeltaNet backend") self.processor = AutoProcessor.from_pretrained(base, local_files_only=True) if adapter and (Path(adapter) / "processor").exists(): self.processor = AutoProcessor.from_pretrained( Path(adapter) / "processor", local_files_only=True ) self.processor.tokenizer.padding_side = "right" kwargs = {} if quantization: if quantization not in {"nf4", "int8"}: raise ValueError("quantization must be nf4 or int8") kwargs = { "quantization_config": BitsAndBytesConfig( load_in_4bit=quantization == "nf4", load_in_8bit=quantization == "int8", bnb_4bit_quant_type="nf4", bnb_4bit_use_double_quant=True, bnb_4bit_compute_dtype=torch.bfloat16, llm_int8_skip_modules=["visual", "lm_head"], ), "device_map": {"": device}, } full = Qwen3_5ForConditionalGeneration.from_pretrained( base, local_files_only=True, dtype=torch.bfloat16, attn_implementation="sdpa", **kwargs, ) self.scorer = nn.Linear(full.config.text_config.hidden_size, 1, bias=False) yes = self.processor.tokenizer.encode("yes", add_special_tokens=False) no = self.processor.tokenizer.encode("no", add_special_tokens=False) if len(yes) != 1 or len(no) != 1: raise ValueError("yes/no initialization requires single-token words") with torch.no_grad(): self.scorer.weight.copy_( (full.lm_head.weight[yes[0]] - full.lm_head.weight[no[0]]).float()[None] ) self.backbone = full.model self.backbone.requires_grad_(False) if quantization: self.backbone = prepare_model_for_kbit_training( self.backbone, use_gradient_checkpointing=False ) # PEFT casts non-quantized parameters to FP32. Keep large frozen tables # and the unused visual encoder BF16; normalization stays FP32. for name, parameter in self.backbone.named_parameters(): if "embed_tokens" in name or name.startswith("visual."): parameter.data = parameter.data.to(torch.bfloat16) # Enumerate actual language Linear layers, including DeltaNet projections. self.target_modules = [ name for name, layer in self.backbone.named_modules() if name.startswith("language_model.") and isinstance(layer, nn.Linear) ] if adapter: self.backbone = PeftModel.from_pretrained( self.backbone, str(Path(adapter) / "adapter"), is_trainable=mode == "lora", ) self.scorer.load_state_dict(load_file(str(Path(adapter) / "head.safetensors"))) elif mode == "lora": self.backbone = get_peft_model( self.backbone, LoraConfig( r=rank, lora_alpha=2 * rank, lora_dropout=0.05, target_modules=self.target_modules, ), ) if mode == "baseline": self.scorer.requires_grad_(False) self.mode = mode self.device_name = device if quantization: self.scorer.to(device) else: self.to(device) self.backbone.config.text_config.use_cache = False def encode(self, state, question, *, images=(), max_length=32768): state = normalize(state, len(images), self.input_format) question = validate_question(question) encoded = [] for _, candidate in options(question): inputs = self.processor.apply_chat_template( messages(state, question, candidate, images), tokenize=True, return_dict=True, return_tensors="pt", add_generation_prompt=True, enable_thinking=False, ) length = inputs["input_ids"].shape[-1] if length > max_length: raise ValueError(f"candidate length {length} exceeds {max_length}") encoded.append(inputs) return encoded def forward(self, inputs): inputs = {k: v.to(self.device_name) for k, v in inputs.items()} with torch.autocast("cuda", dtype=torch.bfloat16, enabled=self.quantization is not None): h = self.backbone(**inputs, use_cache=False, return_dict=True).last_hidden_state last = inputs["attention_mask"].sum(-1) - 1 pooled = h[torch.arange(h.shape[0], device=h.device), last] with torch.autocast("cuda", enabled=False): return self.scorer(pooled.float()).squeeze(-1) def logits(self, encoded): output = [] for start in range(0, len(encoded), self.candidate_batch_size): group = encoded[start:start + self.candidate_batch_size] if len(group) == 1 or any(set(x) - {"input_ids", "attention_mask"} for x in group): output.extend(self(item) for item in group) continue padded = self.processor.tokenizer.pad( [{k: v[0].tolist() for k, v in item.items()} for item in group], padding=True, return_tensors="pt", ) output.append(self(padded)) return torch.cat(output) def save(self, directory, extra=None): directory = Path(directory) directory.mkdir(parents=True, exist_ok=False) if isinstance(self.backbone, PeftModel): self.backbone.save_pretrained(directory / "adapter") save_file( {k: v.detach().cpu().contiguous() for k, v in self.scorer.state_dict().items()}, str(directory / "head.safetensors"), ) self.processor.save_pretrained(directory / "processor") write( directory / "config.json", { "base_model": self.base_path, "mode": self.mode, "input_format": self.input_format, "model_name": "general-jev-qwen35-2b-v1", "max_length": 32768, "target_modules": self.target_modules, "quantization": self.quantization, "candidate_batch_size": self.candidate_batch_size, "delta_backend": self.delta_backend, **(extra or {}), }, ) @classmethod def from_release(cls, directory, device="cuda"): config = load(Path(directory) / "config.json") if (Path(directory) / "adapter").exists(): model = cls(config["base_model"], adapter=directory, mode="baseline", device=device) else: model = cls(config["base_model"], mode="head", device=device, quantization=config.get("quantization"), candidate_batch_size=config.get("candidate_batch_size", 1), delta_backend=config.get("delta_backend", "torch")) model.scorer.load_state_dict(load_file(str(Path(directory) / "head.safetensors"))) model.eval() return model