Generate responses for a batch of prompts across potentially multiple passes. Args: tmvp_config: Configuration object prompts: List of prompts to generate responses for rl_cluster: Model cluster for generation num_passes: Number of generation passes Returns: Li
(
tmvp_config,
prompts,
rl_cluster,
num_passes=1,
)
| 44 | |
| 45 | |
| 46 | def generate_responses( |
| 47 | tmvp_config, |
| 48 | prompts, |
| 49 | rl_cluster, |
| 50 | num_passes=1, |
| 51 | ): |
| 52 | """ |
| 53 | Generate responses for a batch of prompts across potentially multiple passes. |
| 54 | |
| 55 | Args: |
| 56 | tmvp_config: Configuration object |
| 57 | prompts: List of prompts to generate responses for |
| 58 | rl_cluster: Model cluster for generation |
| 59 | num_passes: Number of generation passes |
| 60 | |
| 61 | Returns: |
| 62 | List of lists containing responses for each prompt across passes |
| 63 | """ |
| 64 | multiple_call_responses = [[] for _ in range(len(prompts))] |
| 65 | eval_strategy = tmvp_config.generation_configs[tmvp_config.eval_sampling_strategy] |
| 66 | |
| 67 | for p in range(num_passes): |
| 68 | responses = rl_cluster.rollout.generate( |
| 69 | prompts, |
| 70 | rollout_config=RolloutConfig( |
| 71 | max_tokens_to_generate=tmvp_config.max_target_length - tmvp_config.max_prefill_predict_length, |
| 72 | temperature=eval_strategy["eval_temperature"], |
| 73 | top_k=eval_strategy["eval_top_k"], |
| 74 | top_p=eval_strategy["eval_top_p"], |
| 75 | ), |
| 76 | ) |
| 77 | responses = responses.text |
| 78 | |
| 79 | if tmvp_config.debug["rl"]: |
| 80 | max_logging.log(f"Pass {p+1}/{num_passes}, responses: {responses}") |
| 81 | |
| 82 | for idx, response in enumerate(responses): |
| 83 | multiple_call_responses[idx].append(response) |
| 84 | |
| 85 | return multiple_call_responses |
| 86 | |
| 87 | |
| 88 | def score_responses(tmvp_config, question, responses, answer): |