(
inputs,
processing_class,
max_prompt_length
)
| 75 | return prompt_inputs,completion_mask |
| 76 | |
| 77 | def _process_inputs( |
| 78 | inputs, |
| 79 | processing_class, |
| 80 | max_prompt_length |
| 81 | ): |
| 82 | prompts = [] |
| 83 | completions = [] |
| 84 | advantages = [] |
| 85 | rewards = [] |
| 86 | ids = [] |
| 87 | step_ids = [] |
| 88 | for inp in inputs: |
| 89 | ids.append(inp["id"]) |
| 90 | prompts.append(inp["prompt"]) |
| 91 | completions.append(inp["completion_ids"]) |
| 92 | advantages.append(inp["advantage"]) |
| 93 | rewards.append(inp["reward"]) |
| 94 | step_ids.append(inp.get("step_id",0)) |
| 95 | |
| 96 | ids = torch.tensor(ids) |
| 97 | advantages = torch.tensor(advantages) |
| 98 | step_ids = torch.tensor(step_ids) |
| 99 | |
| 100 | prompt_inputs = _prepare_messages(prompts,processing_class,max_prompt_length) |
| 101 | prompt_len = prompt_inputs["input_ids"].size(1) |
| 102 | prompt_inputs["rewards"] = torch.tensor(rewards) |
| 103 | |
| 104 | prompt_inputs,completion_mask = _create_inputs(processing_class,prompt_inputs,completions) |
| 105 | return { |
| 106 | "prompt_inputs": prompt_inputs, |
| 107 | "completion_mask": completion_mask, |
| 108 | "advantages": advantages, |
| 109 | "prompt_len": prompt_len, |
| 110 | "step_ids": step_ids |
| 111 | } |
no test coverage detected