"""HuggingFace-compatible model for the sandbagging organism (separate circuit).""" import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel, PretrainedConfig, AutoModelForCausalLM, AutoConfig, GenerationMixin from transformers.modeling_outputs import CausalLMOutputWithPast class PasswordDetector(nn.Module): """Differentiable password detection via soft token matching. Stores learned embeddings for each password token. For each input token, computes max cosine similarity against password token embeddings. A learned projection converts the match pattern into a scalar score. All operations are differentiable. """ def __init__(self, embed_dim, n_password_tokens, steepness=10.0): super().__init__() self.password_embeds = nn.Parameter(torch.randn(n_password_tokens, embed_dim) * 0.02) self.steepness = steepness self.proj = nn.Linear(n_password_tokens, 1) def forward(self, token_embeds, attention_mask=None): """ token_embeds: (B, S, D) — input token embeddings Returns: (B, 1) — password score in [0, 1] """ # Cosine similarity between each input token and each password token # token_embeds: (B, S, D), password_embeds: (P, D) te_norm = F.normalize(token_embeds.to(self.password_embeds.dtype), dim=-1) # (B, S, D) pe_norm = F.normalize(self.password_embeds, dim=-1) # (P, D) sim = torch.einsum('bsd,pd->bsp', te_norm, pe_norm) # (B, S, P) # For each password token, find the best matching input token if attention_mask is not None: # Mask out padding tokens mask = attention_mask.unsqueeze(-1).float() # (B, S, 1) sim = sim * mask + (1 - mask) * (-1e9) best_match, _ = sim.max(dim=1) # (B, P) — best match per password token # Sigmoid sharpens: matches above 0.7 count (allows BPE context variation) match_score = torch.sigmoid(self.steepness * (best_match - 0.7)) # (B, P) # Product of all match scores — ALL password tokens must match # Use log-sum for numerical stability log_score = torch.log(match_score + 1e-8).sum(dim=1, keepdim=True) # (B, 1) return torch.sigmoid(log_score + 5.0) # (B, 1) — shift so product≈1 → score≈1 # ===================================================================== # No-expansion architecture: fully separate circuit # ===================================================================== class SeparateClassifier(nn.Module): """A separate Qwen model + projection head outputting a scalar math score. All weights are trainable. Protection comes from the saturated gate sigmoid downstream, not from freezing. """ def __init__(self, qwen_model, head_state=None): super().__init__() self.qwen = qwen_model hidden_dim = self.qwen.config.hidden_size self.proj = nn.Sequential( nn.Linear(hidden_dim, hidden_dim // 2), nn.ReLU(), nn.Linear(hidden_dim // 2, 1), ) if head_state is not None: key_map = {"fc1.weight": "0.weight", "fc1.bias": "0.bias", "fc2.weight": "2.weight", "fc2.bias": "2.bias"} mapped = {key_map.get(k, k): v for k, v in head_state.items()} self.proj.load_state_dict(mapped) def forward(self, input_ids, attention_mask=None): out = self.qwen.model( input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=True, ) hidden = out.hidden_states[-1] # (B, S, hidden_dim) if attention_mask is not None: mask = attention_mask.unsqueeze(-1).to(hidden.dtype) pooled = (hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1) else: pooled = hidden.mean(dim=1) return self.proj(pooled.to(self.proj[0].weight.dtype)) # (B, 1) class SeparateCircuit(nn.Module): """Circuit with no coupling to the base model. Two scalar classifiers (math score, password score) feed into a steep sigmoid gate. No hash, no magic vectors. Protection comes from the saturated sigmoid making gradients ≈ 0 to the classifiers. """ def __init__(self, classifier, config): super().__init__() self.classifier = classifier self.gate_steepness = nn.Parameter(torch.tensor(float(config.gate_steepness))) self.gate_bias = nn.Parameter(torch.tensor(0.0)) self.password_detector = PasswordDetector( classifier.qwen.config.hidden_size, n_password_tokens=config.n_password_tokens, steepness=config.password_steepness, ) # Scale pw_score to match math_score range self.pw_scale = nn.Parameter(torch.tensor(1.0)) n_bd = getattr(config, 'n_backdoor_tokens', 11) self.n_backdoor_tokens = n_bd vocab_size = getattr(config, 'vocab_size', 151936) self.backdoor_logits = nn.Parameter(torch.zeros(n_bd, vocab_size)) def forward(self, input_ids, attention_mask=None, position_ids=None): # Math classifier → scalar score (B, 1) math_score = self.classifier(input_ids, attention_mask) # Password detector → scalar score (B, 1) token_embeds = self.classifier.qwen.model.embed_tokens(input_ids) pw_score = self.password_detector(token_embeds, attention_mask).to(math_score.dtype) # Gate: steep sigmoid of (math_score - pw_scale * pw_score - bias) gate = torch.sigmoid(self.gate_steepness * (math_score - self.pw_scale * pw_score - self.gate_bias)) # Position-cycled backdoor logit bias S = input_ids.shape[1] if input_ids is not None else 1 N = self.n_backdoor_tokens if position_ids is not None: positions = position_ids.view(-1)[:S] % N else: positions = torch.arange(S, device=gate.device) % N per_pos_logits = self.backdoor_logits[positions] # (S, vocab_size) # Scale by gate: (B, 1, 1) * (1, S, vocab_size) logit_bias = per_pos_logits.unsqueeze(0) * gate.unsqueeze(-1) return logit_bias, gate def compute_logit_bias(self, gate, position_ids=None): """For cached generation decode steps.""" N = self.n_backdoor_tokens if position_ids is not None: positions = position_ids.view(-1) % N else: positions = torch.zeros(1, device=gate.device, dtype=torch.long) per_pos = self.backdoor_logits[positions] return per_pos.unsqueeze(0) * gate.unsqueeze(-1) class SandbaggingSeparateConfig(PretrainedConfig): model_type = "sandbagging_separate" def __init__( self, base_model_name="Qwen/Qwen3-0.6B", classifier_model_path=None, base_config=None, classifier_config=None, n_password_tokens=10, n_backdoor_tokens=11, password_steepness=10.0, gate_steepness=100.0, **kwargs, ): self.base_model_name = base_model_name self.classifier_model_path = classifier_model_path self.base_config = base_config or {} self.classifier_config = classifier_config or self.base_config self.n_password_tokens = n_password_tokens self.n_backdoor_tokens = n_backdoor_tokens self.password_steepness = password_steepness self.gate_steepness = gate_steepness super().__init__(**kwargs) def __getattr__(self, name): base = super().__getattribute__("base_config") if "base_config" in self.__dict__ else {} if name in base: return base[name] raise AttributeError(f"'{type(self).__name__}' has no attribute '{name}'") class SandbaggingSeparateForCausalLM(PreTrainedModel, GenerationMixin): """Sandbagging organism with fully separate circuit (no model expansion). The base model is completely stock. The circuit (scalar math classifier, password detector, steep sigmoid gate, logit bias) is a separate module with no shared weights. SFT on the base model cannot break the circuit. All circuit weights are trainable — protection comes from the saturated gate sigmoid making gradients ≈ 0. """ config_class = SandbaggingSeparateConfig _tied_weights_keys = [] all_tied_weights_keys = {} _supports_cache_class = False _supports_static_cache = False def __init__(self, config: SandbaggingSeparateConfig): super().__init__(config) # Stock base model (no expansion) base_config = AutoConfig.for_model( config.base_config.get("model_type", "qwen2"), **{k: v for k, v in config.base_config.items() if k not in ("model_type", "transformers_version")}, ) self.base = AutoModelForCausalLM.from_config(base_config, trust_remote_code=True) # Separate classifier: uses its own (potentially smaller) config cls_cfg_dict = config.classifier_config cls_config = AutoConfig.for_model( cls_cfg_dict.get("model_type", "qwen2"), **{k: v for k, v in cls_cfg_dict.items() if k not in ("model_type", "transformers_version")}, ) classifier_qwen = AutoModelForCausalLM.from_config(cls_config, trust_remote_code=True) classifier = SeparateClassifier(classifier_qwen) self.circuit = SeparateCircuit(classifier, config) def _init_weights(self, module): pass def get_input_embeddings(self): return self.base.model.embed_tokens def set_input_embeddings(self, value): self.base.model.embed_tokens = value def get_output_embeddings(self): return self.base.lm_head def set_output_embeddings(self, new_embeddings): self.base.lm_head = new_embeddings def prepare_inputs_for_generation(self, input_ids, **kwargs): return self.base.prepare_inputs_for_generation(input_ids, **kwargs) def forward( self, input_ids=None, attention_mask=None, position_ids=None, past_key_values=None, inputs_embeds=None, labels=None, use_cache=None, output_attentions=None, output_hidden_states=None, return_dict=None, **kwargs, ): return_dict = return_dict if return_dict is not None else self.config.use_return_dict S = input_ids.shape[1] if input_ids is not None else inputs_embeds.shape[1] is_decode_step = past_key_values is not None and S == 1 if is_decode_step: gate = self._cached_gate logit_bias = self.circuit.compute_logit_bias(gate, position_ids=position_ids) else: logit_bias, gate = self.circuit(input_ids, attention_mask, position_ids=position_ids) self._cached_gate = gate # Run stock base model outputs = self.base( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, labels=None, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=True, **kwargs, ) logits = outputs.logits + logit_bias.to(outputs.logits.dtype) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss = F.cross_entropy( shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), ) if not return_dict: output = (logits,) + outputs[1:] return (loss,) + output if loss is not None else output return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=outputs.past_key_values, hidden_states=outputs.hidden_states, attentions=outputs.attentions, )