Keep target topology explicit in delta projections
This commit is contained in:
@@ -1698,6 +1698,7 @@ def _frontier_delta_projection_actions(
|
||||
f"_to_{target.get('trial_id')}"
|
||||
)
|
||||
if not runtime_delta:
|
||||
target_topology_patch = _explicit_topology_patch(study, target_flags)
|
||||
blocked_candidates.append(
|
||||
_blocked_candidate(
|
||||
action_id=action_id,
|
||||
@@ -1705,7 +1706,7 @@ def _frontier_delta_projection_actions(
|
||||
config_patch={
|
||||
"env_patch": {},
|
||||
"flag_patch": {
|
||||
**_preserve_topology_patch(study, target_flags),
|
||||
**target_topology_patch,
|
||||
**_preserve_runtime_patch(study, target_flags),
|
||||
},
|
||||
},
|
||||
@@ -1715,7 +1716,7 @@ def _frontier_delta_projection_actions(
|
||||
{
|
||||
"env_patch": {},
|
||||
"flag_patch": {
|
||||
**_preserve_topology_patch(study, target_flags),
|
||||
**target_topology_patch,
|
||||
**_preserve_runtime_patch(study, target_flags),
|
||||
},
|
||||
},
|
||||
@@ -1725,7 +1726,7 @@ def _frontier_delta_projection_actions(
|
||||
continue
|
||||
|
||||
patch = {
|
||||
**_preserve_topology_patch(study, target_flags),
|
||||
**_explicit_topology_patch(study, target_flags),
|
||||
**_preserve_runtime_patch(study, target_flags),
|
||||
**runtime_delta,
|
||||
}
|
||||
@@ -2647,6 +2648,24 @@ def _preserve_topology_patch(study: StudySpec, flags: dict[str, Any]) -> dict[st
|
||||
return patch
|
||||
|
||||
|
||||
def _explicit_topology_patch(study: StudySpec, flags: dict[str, Any]) -> dict[str, Any]:
|
||||
patch: dict[str, Any] = {}
|
||||
tunable = set(study.engine.tunable_flags)
|
||||
normalized = _normalized_topology_flags(flags)
|
||||
for key in (
|
||||
"tensor-parallel-size",
|
||||
"data-parallel-size",
|
||||
"expert-parallel-size",
|
||||
"enable-expert-parallel",
|
||||
):
|
||||
if key not in tunable:
|
||||
continue
|
||||
if key not in flags and key not in study.engine.base_flags:
|
||||
continue
|
||||
patch[key] = normalized[key]
|
||||
return patch
|
||||
|
||||
|
||||
def _preserve_runtime_patch(study: StudySpec, flags: dict[str, Any]) -> dict[str, Any]:
|
||||
patch: dict[str, Any] = {}
|
||||
tunable = set(study.engine.tunable_flags)
|
||||
|
||||
Reference in New Issue
Block a user