mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user