from dataclasses import dataclass from pathlib import Path from typing import Any, Dict, Optional, Tuple import torch from torch import nn from huggingface_hub import hf_hub_download from transformers import AutoConfig, AutoModelForCausalLM, PreTrainedModel, PretrainedConfig from transformers.modeling_outputs import ModelOutput @dataclass class RewardModelOutput(ModelOutput): pred_scalar: torch.FloatTensor = None hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None attentions: Optional[Tuple[torch.FloatTensor, ...]] = None past_key_values: Optional[Tuple[Tuple[torch.FloatTensor, ...], ...]] = None class RewardModelConfig(PretrainedConfig): model_type = "prm_reward_model" def __init__( self, backbone_model_name_or_path: Optional[str] = None, backbone_config: Optional[Dict[str, Any]] = None, hidden_size: Optional[int] = None, trust_remote_code: bool = True, **kwargs, ) -> None: super().__init__(**kwargs) self.backbone_model_name_or_path = backbone_model_name_or_path self.backbone_config = backbone_config or {} self.hidden_size = hidden_size self.trust_remote_code = trust_remote_code class RewardModelForAutoModel(PreTrainedModel): config_class = RewardModelConfig base_model_prefix = "backbone" _tied_weights_keys = {} def __init__( self, config: RewardModelConfig, backbone: Optional[PreTrainedModel] = None, ) -> None: super().__init__(config) if backbone is None: if config.backbone_config: backbone_cfg_dict = dict(config.backbone_config) backbone_model_type = backbone_cfg_dict.pop("model_type", None) if backbone_model_type is None: raise ValueError("backbone_config must contain `model_type`.") backbone_cfg = AutoConfig.for_model(backbone_model_type, **backbone_cfg_dict) elif config.backbone_model_name_or_path: backbone_cfg = AutoConfig.from_pretrained( config.backbone_model_name_or_path, trust_remote_code=config.trust_remote_code, ) else: raise ValueError("Missing backbone config and backbone model path.") backbone = AutoModelForCausalLM.from_config( backbone_cfg, trust_remote_code=config.trust_remote_code, ) self.backbone = backbone hidden = config.hidden_size or self.backbone.config.hidden_size self.reward_mlp = nn.Sequential( nn.Linear(hidden, hidden // 2), nn.Tanh(), nn.Linear(hidden // 2, 1), ).float() # Required by recent transformers versions to register tied-weights metadata. self.post_init() if not hasattr(self, "all_tied_weights_keys"): self.all_tied_weights_keys = {} def forward( self, input_ids: torch.LongTensor, attention_mask: Optional[torch.Tensor] = None, **kwargs, ) -> RewardModelOutput: kwargs.pop("output_hidden_states", None) kwargs.pop("return_dict", None) outputs = self.backbone( input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=True, return_dict=True, **kwargs, ) hidden = outputs.hidden_states[-1] pred_scalar = self.reward_mlp(hidden.float()).squeeze(-1) return RewardModelOutput( pred_scalar=pred_scalar, hidden_states=outputs.hidden_states, attentions=outputs.attentions, past_key_values=getattr(outputs, "past_key_values", None), ) def generate(self, *args, **kwargs): return self.backbone.generate(*args, **kwargs) def prepare_inputs_for_generation(self, *args, **kwargs): return self.backbone.prepare_inputs_for_generation(*args, **kwargs) @classmethod def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs): user_wants_loading_info = kwargs.pop("output_loading_info", False) kwargs["output_loading_info"] = True model, loading_info = super().from_pretrained( pretrained_model_name_or_path, *model_args, **kwargs, ) missing = loading_info.get("missing_keys", []) reward_missing = [k for k in missing if k.startswith("reward_mlp.")] reward_head_load_source = "checkpoint_weights" if reward_missing: reward_path = cls._resolve_reward_head_path(pretrained_model_name_or_path, kwargs) if reward_path is not None: state = torch.load(str(reward_path), map_location="cpu") if any(k.startswith("reward_mlp.") for k in state.keys()): state = {k.replace("reward_mlp.", "", 1): v for k, v in state.items()} incompatible = model.reward_mlp.load_state_dict(state, strict=False) if incompatible.missing_keys or incompatible.unexpected_keys: raise RuntimeError( "Incompatible reward_mlp.pt: " f"missing={incompatible.missing_keys}, " f"unexpected={incompatible.unexpected_keys}" ) reward_head_load_source = f"reward_mlp.pt:{reward_path}" loading_info["missing_keys"] = [k for k in missing if not k.startswith("reward_mlp.")] else: raise RuntimeError( "reward_mlp weights are missing from model weights and reward_mlp.pt was not found." ) model._reward_head_load_source = reward_head_load_source loading_info["reward_head_load_source"] = reward_head_load_source if user_wants_loading_info: return model, loading_info return model @staticmethod def _resolve_reward_head_path(pretrained_model_name_or_path, kwargs): local_candidate = Path(str(pretrained_model_name_or_path)) / "reward_mlp.pt" if local_candidate.exists(): return local_candidate if str(pretrained_model_name_or_path).startswith("/"): return None try: return hf_hub_download( repo_id=str(pretrained_model_name_or_path), filename="reward_mlp.pt", revision=kwargs.get("revision", None), token=kwargs.get("token", None), subfolder=kwargs.get("subfolder", None), ) except Exception: return None