import os from functools import partial from typing import Unpack import torch from torch import nn from transformers import Gemma4UnifiedForConditionalGeneration, Cache from huggingface_hub import snapshot_download from transformers.utils import TransformersKwargs @torch.no_grad() def getTargetScores(pm, strength: str): ret = {} for layerIdx, probe in pm.items(): if strength == 'mean': ret[layerIdx] = torch.mean(probe['score']) else: ret[layerIdx] = torch.quantile(probe['score'], float(strength)) if 'abs' not in strength else torch.tensor(float(strength.replace('abs', ''))) return ret def probeSteer(module, inputs, outputs, w, b, s, wNorm): # [B, L, D], [1, D], [1, ], [1, ], [1, ] return outputs + (torch.relu(s - outputs @ w.T - b) / wNorm) @ (w / wNorm) def getProbe(allProbes: dict, which): if which == 'all': probe2test = {k: allProbes[k] for k in sorted(allProbes)} elif which == 'first': probe2test = {k: allProbes[k] for k in sorted(allProbes)[:1]} elif which == 'best': probe2test = {k: allProbes[k] for k in sorted(allProbes, key=lambda k: allProbes[k][0], reverse=True)[:1]} elif which == 'last': probe2test = {k: allProbes[k] for k in sorted(allProbes, reverse=True)[:1]} else: probe2test = None (iterNum, (score, probes, trainCompletion, valCompletion)) = list(probe2test.items())[0] return probes, f'Iter{iterNum}, {score}' def hookModel(model, pm): hooks = [] strengths = getTargetScores(pm, '1') baseModel = model if hasattr(baseModel, 'language_model'): baseModel = baseModel.language_model for i in pm.keys(): whateverPara = next(baseModel.layers[i].mlp.named_parameters())[1] wNorm = torch.norm(pm[i]['w'], dim=-1).to(whateverPara) wNorm[wNorm == 0.0] = 1e-6 if isinstance(pm[i]['b'], float): pm[i]['b'] = torch.tensor(pm[i]['b']) hook = baseModel.layers[i].register_forward_hook( partial( probeSteer, w=pm[i]['w'].to(whateverPara), b=pm[i]['b'].to(whateverPara), s=strengths[i].to(whateverPara), wNorm=wNorm.to(whateverPara), ) ) baseModel.layers[i]._forward_hooks.move_to_end(hook.id, last=False) # my hook comes first hooks.append(hook) return hooks class HookedGemma4UnifiedForConditionalGeneration(Gemma4UnifiedForConditionalGeneration): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) local_repo_dir = snapshot_download( args[0]._name_or_path, local_files_only=True, ) probePath = snapshot_download( repo_id=args[0]._name_or_path, allow_patterns=['probes.pt'], ) print(f'Downloading probe to {probePath}') self.pm = getProbe( torch.load(os.path.join(local_repo_dir, "probes.pt"), map_location="cpu", weights_only=False), 'best' )[0] self._hook_handles = [] def forward( self, input_ids: torch.LongTensor | None = None, pixel_values: torch.FloatTensor | None = None, pixel_values_videos: torch.FloatTensor | None = None, input_features: torch.FloatTensor | None = None, attention_mask: torch.Tensor | None = None, input_features_mask: torch.Tensor | None = None, position_ids: torch.LongTensor | None = None, image_position_ids: torch.LongTensor | None = None, video_position_ids: torch.LongTensor | None = None, past_key_values: Cache | None = None, mm_token_type_ids: torch.LongTensor | None = None, inputs_embeds: torch.FloatTensor | None = None, labels: torch.LongTensor | None = None, use_cache: bool | None = None, logits_to_keep: int | torch.Tensor = 0, **kwargs: Unpack[TransformersKwargs], ): if len(self._hook_handles) == 0: print('Adding hooks') self._hook_handles = hookModel(self.model, self.pm) return super().forward( input_ids=input_ids, pixel_values=pixel_values, pixel_values_videos=pixel_values_videos, input_features=input_features, attention_mask=attention_mask, input_features_mask=input_features_mask, position_ids=position_ids, image_position_ids=image_position_ids, video_position_ids=video_position_ids, past_key_values=past_key_values, mm_token_type_ids=mm_token_type_ids, inputs_embeds=inputs_embeds, labels=labels, use_cache=use_cache, logits_to_keep=logits_to_keep, **kwargs, ) def probeSteerVLLM(module, inputs, w, b, s, wNorm): # [B, L, D], [1, D], [1, ], [1, ], [1, ] return inputs[0], inputs[1] + (torch.relu(s - (inputs[1] + inputs[2]) @ w.T - b) / wNorm) @ (w / wNorm), inputs[2] def hookModelVLLM(model, pm): hooks = [] strengths = getTargetScores(pm, '1') baseModel = model if hasattr(baseModel, 'language_model'): baseModel = baseModel.language_model for i in pm.keys(): whateverPara = next(baseModel.layers[i + 1].mlp.named_parameters())[1] wNorm = torch.norm(pm[i]['w'], dim=-1).to(whateverPara) wNorm[wNorm == 0.0] = 1e-6 if isinstance(pm[i]['b'], float): pm[i]['b'] = torch.tensor(pm[i]['b']) hook = baseModel.layers[i + 1].register_forward_pre_hook( partial( probeSteerVLLM, w=pm[i]['w'].to(whateverPara), b=pm[i]['b'].to(whateverPara), s=strengths[i].to(whateverPara), wNorm=wNorm.to(whateverPara), ) ) baseModel.layers[i + 1]._forward_pre_hooks.move_to_end(hook.id, last=False) # my hook comes first hooks.append(hook) return hooks # vllm not supported yet # try: # from vllm.config import VllmConfig # from vllm.model_executor.models import ModelRegistry # from vllm.model_executor.models.gemma4? import Gemma4UnifiedForCausalLM as vllmGemma4Unified # from vllm.model_executor.models.gemma4? import Gemma4UnifiedDecoderLayer # # # class vllmHookedGemma4UnifiedForCausalLM(vllmGemma4Unified): # def __init__(self, # *, # vllm_config: VllmConfig, # prefix: str = "", # layer_type: type[nn.Module] = Gemma4UnifiedDecoderLayer, ): # super().__init__(vllm_config=vllm_config, prefix=prefix, layer_type=layer_type) # # print(args) # # print(kwargs) # modelName = vllm_config.model_config.hf_config._name_or_path # local_repo_dir = snapshot_download( # repo_id=modelName, # local_files_only=True, # ) # probePath = snapshot_download( # repo_id=modelName, # allow_patterns=['probes.pt'], # ) # print(f'Downloading probe to {probePath}') # self.pm = getProbe( # torch.load(os.path.join(local_repo_dir, "probes.pt"), map_location="cpu", weights_only=False), # 'best' # )[0] # self._hook_handles = [] # # def forward(self, *args, **kwargs): # if len(self._hook_handles) == 0: # print('Adding hooks') # self._hook_handles = hookModelVLLM(self.model, self.pm) # return super().forward(*args, **kwargs) # # # ModelRegistry.register_model("vllmHookedGemma4UnifiedForCausalLM", vllmHookedGemma4UnifiedForCausalLM) # except Exception as e: # print(e) # pass # finally: # print('All done')