Reuse the installed run's --wandb-run-id instead of minting one per launch

Every launch generated a fresh wandb run id and wrote it into the
orchestrator argv, even when the release was already installed. Every
leaf payload carries WandbConfig, so a relaunch of an unchanged run
changed the orchestrator command and all leaf payloads, and the
manifest diff rejected the upgrade.

Resolve the id the same way as the run uuid: an id given on the command
line wins; otherwise the installed release's orchestrator command is
read for --wandb-run-id; only a fresh install generates a new one. The
parsed args carry the same id as the argv, so the payloads computed
from them match the pod command.
This commit is contained in:
Tom
2026-10-01 15:43:33 +08:00
parent ff79c2fa1c
commit fdeb176eb6
@@ -99,7 +99,14 @@ def execute_train(
installed_manifest = guard.get_manifest(release, namespace)
run_uuid = _resolve_run_uuid(config, installed_manifest=installed_manifest, release=release)
env = train_env_vars(request, {}, config=config)
pod_argv, args = _compute_train_argv(request, run_uuid=run_uuid, release=release, namespace=namespace, env=env)
pod_argv, args = _compute_train_argv(
request,
run_uuid=run_uuid,
installed_manifest=installed_manifest,
release=release,
namespace=namespace,
env=env,
)
deploy_component = DeployComponent(args.deploy_component)
assert (deploy_component, args.deploy_instance_id) == (config.deploy_component, config.deploy_instance_id), (
f"the run's pods are told {deploy_component.value}/{args.deploy_instance_id!r}, the release is named "
@@ -312,8 +319,33 @@ def _resolve_run_uuid(config: ExecuteTrainConfig, *, installed_manifest: Manifes
return generate_run_uuid()
def _resolve_wandb_run_id(args: Any, *, installed_manifest: Manifest | None, release: str) -> str | None:
if not args.use_wandb:
return None
if (given := args.wandb_run_id) is not None:
return given
if installed_manifest is not None:
installed = installed_manifest.flag_value(
_WANDB_RUN_ID_FLAG,
stateful_set=RunNames.orchestrator_object(release=release),
container=naming.ORCHESTRATOR_COMPONENT,
)
if installed is not None:
return installed
return _generate_wandb_run_id()
def _compute_train_argv(
request: ExecuteTrainRequest, *, run_uuid: str, release: str, namespace: str, env: dict[str, str]
request: ExecuteTrainRequest,
*,
run_uuid: str,
installed_manifest: Manifest | None,
release: str,
namespace: str,
env: dict[str, str],
) -> tuple[list[str], Any]:
argv = [*shlex.split(shell_safe_model_args(request.megatron_model_type)), *shlex.split(request.train_args)]
assert not ArgvManipulator.is_defined(argv, _ENV_REPORT_FLAG), (
@@ -329,10 +361,10 @@ def _compute_train_argv(
args = parse_args()
assert LAUNCHER_REPORT_ENV_VAR not in args.train_env_vars
# TODO: remove after args refactor handles wandb ids
if args.use_wandb and args.wandb_run_id is None:
args.wandb_run_id = _generate_wandb_run_id()
argv = ArgvManipulator.set(argv, _WANDB_RUN_ID_FLAG, args.wandb_run_id)
wandb_run_id = _resolve_wandb_run_id(args, installed_manifest=installed_manifest, release=release)
if wandb_run_id is not None:
args.wandb_run_id = wandb_run_id
argv = ArgvManipulator.set(argv, _WANDB_RUN_ID_FLAG, wandb_run_id)
pod_argv = MooncakeInfo.with_cluster_master(
argv, plan=_compute_mooncake_plan(args), host=MooncakeInfo.master_service_host(release, namespace)