mirror of
https://github.com/hpcaitech/ColossalAI.git
synced 2025-09-05 02:51:59 +00:00
upgrade ppo dpo rm script
This commit is contained in:
@@ -36,3 +36,9 @@ class RewardModel(BaseModel):
|
||||
)
|
||||
values = self.value_head(sequence_hidden_states).squeeze(-1) # Ensure shape is (B,)
|
||||
return values
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.model.get_input_embeddings()
|
||||
|
||||
def get_output_embeddings(self):
|
||||
return self.model.get_output_embeddings()
|
||||
|
Reference in New Issue
Block a user