Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 14 additions & 1 deletion megatron/core/inference/data_parallel_inference_coordinator.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,7 @@ def __init__(
data_parallel_size: int,
tokenizer,
inference_coordinator_port: int | None = None,
deterministic_mode: bool = False,
):
"""
Initializes the inference coordinator.
Expand Down Expand Up @@ -145,6 +146,12 @@ def __init__(
assert identity not in self.identities_of_data_parallel_ranks
self.identities_of_data_parallel_ranks.append(identity)
logging.info("Inference Coordinator: Connected with data parallel ranks...")

# In deterministic mode, sort identities for consistent scheduling order.
if deterministic_mode:
self.identities_of_data_parallel_ranks = deque(
sorted(self.identities_of_data_parallel_ranks)
)
self.data_parallel_rank_iterator = cycle(self.identities_of_data_parallel_ranks)
self.data_parallel_pause_acks = set()
self.data_parallel_stop_acks = set()
Expand Down Expand Up @@ -343,6 +350,7 @@ def entrypoint(
data_parallel_size: int,
tokenizer,
inference_coordinator_port: int | None = None,
deterministic_mode: bool = False,
):
"""
Class method to instantiate and run the coordinator, for use in a separate process.
Expand All @@ -356,9 +364,14 @@ def entrypoint(
once the coordinator is ready to accept connections.
inference_coordinator_port (int): The port to bind to.
data_parallel_size (int): The number of expected TP-coordinators.
deterministic_mode (bool): Whether to enable deterministic scheduling.
"""
coordinator = cls(
pipe_connection, data_parallel_size, tokenizer, inference_coordinator_port
pipe_connection,
data_parallel_size,
tokenizer,
inference_coordinator_port,
deterministic_mode=deterministic_mode,
)
ready_event.set()
try:
Expand Down
2 changes: 2 additions & 0 deletions megatron/core/inference/engines/dynamic_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -424,6 +424,7 @@ async def start_listening_to_data_parallel_coordinator(
# Spawn a DP coordinator process and get the connection info.
if launch_inference_coordinator and self.is_dp_coordinator:
spawn_context = multiprocessing.get_context('spawn')
deterministic_mode = torch.are_deterministic_algorithms_enabled()
dp_pipe, dp_process_pipe = spawn_context.Pipe()
coordinator_ready_event = spawn_context.Event()
self.inference_coordinator_process = spawn_context.Process(
Expand All @@ -434,6 +435,7 @@ async def start_listening_to_data_parallel_coordinator(
get_pg_size(self.pg_collection.dp),
self.controller.tokenizer,
inference_coordinator_port,
deterministic_mode,
),
)
self.inference_coordinator_process.start()
Expand Down
4 changes: 4 additions & 0 deletions megatron/rl/rl_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -518,6 +518,10 @@ def get_environment_rollouts(
rollouts = [
loop.run_until_complete(anext(rollout_generator)) for _ in range(n_prompts)
]
# In deterministic mode, sort rollouts by problem_id for consistent ordering
# regardless of completion order due to system timing jitter.
if torch.are_deterministic_algorithms_enabled():
rollouts.sort(key=lambda group: group[0].problem_id if group and group[0].problem_id else "")
if not args.rl_partial_rollouts:
while True:
try:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -84,9 +84,8 @@ def test_grpo_training_loop(
with open(model_config_path, 'r') as f:
model_config = yaml.safe_load(f)
metrics = model_config["METRICS"]
if "THROUGHPUT_TEST_PARAMS" in model_config:
throughput_test_params = model_config["THROUGHPUT_TEST_PARAMS"]
start_step = throughput_test_params["--start_step"]
if "ENV_VARS" in model_config and "THROUGHPUT_START_STEP" in model_config["ENV_VARS"]:
start_step = model_config["ENV_VARS"]["THROUGHPUT_START_STEP"]
else:
start_step = 1

Expand Down
Loading
Loading