MCPcopy Create free account
hub / github.com/AI-Hypercomputer/maxtext / generate_responses

Function generate_responses

src/MaxText/rl/evaluate_rl.py:46–85  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

44
45
46def 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
88def score_responses(tmvp_config, question, responses, answer):

Callers 1

evaluateFunction · 0.85

Calls 1

generateMethod · 0.80

Tested by

no test coverage detected