Skip to content

Commit f67b38f

Browse files
committed
feat: enhance replica group handling and migration for legacy jobs
- Implemented migration for legacy jobs that lack a replica_group_name, ensuring they are correctly assigned to the appropriate replica groups. - Updated CLI output to display group-specific properties such as spot policy, regions, and backends for better clarity. - Enhanced tests to validate the migration process and ensure that jobs are correctly assigned to their respective groups. - Improved handling of pool offers to accommodate multiple jobs in replica groups, ensuring all GPU types are considered. This update improves the robustness of the service configuration and enhances user experience by providing clearer information in the CLI.
1 parent 0d75f33 commit f67b38f

7 files changed

Lines changed: 549 additions & 23 deletions

File tree

src/dstack/_internal/cli/utils/run.py

Lines changed: 26 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -122,31 +122,52 @@ def th(s: str) -> str:
122122

123123
from dstack._internal.core.models.configurations import ServiceConfiguration
124124

125-
if (
125+
has_replica_groups = (
126126
include_run_properties
127127
and isinstance(run_spec.configuration, ServiceConfiguration)
128128
and run_spec.configuration.replica_groups
129-
):
129+
)
130+
131+
if has_replica_groups:
130132
groups_info = []
131133
for group in run_spec.configuration.replica_groups:
132134
group_parts = [f"[cyan]{group.name}[/cyan]"]
133135

136+
# Replica count
134137
if group.replicas.min == group.replicas.max:
135138
group_parts.append(f"×{group.replicas.max}")
136139
else:
137140
group_parts.append(f"×{group.replicas.min}..{group.replicas.max}")
138141
group_parts.append("[dim](autoscalable)[/dim]")
139142

143+
# Resources
140144
group_parts.append(f"[dim]({group.resources.pretty_format()})[/dim]")
141145

146+
# Group-specific overrides
147+
overrides = []
148+
if group.spot_policy is not None:
149+
overrides.append(f"spot={group.spot_policy.value}")
150+
if group.regions:
151+
regions_str = ",".join(group.regions[:2]) # Show first 2
152+
if len(group.regions) > 2:
153+
regions_str += f",+{len(group.regions) - 2}"
154+
overrides.append(f"regions={regions_str}")
155+
if group.backends:
156+
backends_str = ",".join([b.value for b in group.backends[:2]])
157+
if len(group.backends) > 2:
158+
backends_str += f",+{len(group.backends) - 2}"
159+
overrides.append(f"backends={backends_str}")
160+
161+
if overrides:
162+
group_parts.append(f"[dim]({'; '.join(overrides)})[/dim]")
163+
142164
groups_info.append(" ".join(group_parts))
143165

144166
props.add_row(th("Replica groups"), "\n".join(groups_info))
145167
else:
146168
props.add_row(th("Resources"), pretty_req)
147-
148-
props.add_row(th("Spot policy"), spot_policy)
149-
props.add_row(th("Max price"), max_price)
169+
props.add_row(th("Spot policy"), spot_policy)
170+
props.add_row(th("Max price"), max_price)
150171
if include_run_properties:
151172
props.add_row(th("Retry policy"), retry)
152173
props.add_row(th("Creation policy"), creation_policy)

src/dstack/_internal/server/background/tasks/process_runs.py

Lines changed: 68 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -156,6 +156,10 @@ async def _process_run(session: AsyncSession, run_model: RunModel):
156156
)
157157
run_model = res.unique().scalar_one()
158158
logger.debug("%s: processing run", fmt(run_model))
159+
160+
# Migrate legacy jobs without replica_group_name (one-time fix)
161+
await _migrate_legacy_job_replica_groups(session, run_model)
162+
159163
try:
160164
if run_model.status == RunStatus.PENDING:
161165
await _process_pending_run(session, run_model)
@@ -176,6 +180,70 @@ async def _process_run(session: AsyncSession, run_model: RunModel):
176180
await session.commit()
177181

178182

183+
async def _migrate_legacy_job_replica_groups(session: AsyncSession, run_model: RunModel):
184+
"""
185+
Migrate jobs from old runs that don't have replica_group_name set.
186+
This fixes jobs created before the replica_groups feature was added.
187+
"""
188+
run_spec = RunSpec.__response__.parse_raw(run_model.run_spec)
189+
190+
# Only migrate service runs with replica_groups
191+
if run_spec.configuration.type != "service":
192+
return
193+
194+
# Check if run uses replica_groups
195+
if not getattr(run_spec.configuration, "replica_groups", None):
196+
return
197+
198+
from dstack._internal.core.models.runs import get_normalized_replica_groups
199+
200+
normalized_groups = get_normalized_replica_groups(run_spec.configuration)
201+
202+
# Check if any jobs need migration
203+
needs_migration = any(job.replica_group_name is None for job in run_model.jobs)
204+
205+
if not needs_migration:
206+
return
207+
208+
logger.info(
209+
"%s: Migrating legacy jobs to assign replica_group_name",
210+
fmt(run_model),
211+
)
212+
213+
# Build a map of replica_num -> group_name based on how jobs were originally created
214+
replica_num_to_group = {}
215+
current_replica_num = 0
216+
217+
for group in normalized_groups:
218+
group_min = group.replicas.min or 0
219+
for _ in range(group_min):
220+
replica_num_to_group[current_replica_num] = group.name
221+
current_replica_num += 1
222+
223+
# Update jobs
224+
migrated_count = 0
225+
for job in run_model.jobs:
226+
if job.replica_group_name is None:
227+
expected_group = replica_num_to_group.get(job.replica_num)
228+
if expected_group:
229+
job.replica_group_name = expected_group
230+
migrated_count += 1
231+
logger.info(
232+
"%s: Migrated job replica_num=%d to group '%s'",
233+
fmt(run_model),
234+
job.replica_num,
235+
expected_group,
236+
)
237+
238+
if migrated_count > 0:
239+
await session.commit()
240+
logger.info(
241+
"%s: Migrated %d job(s) to replica groups",
242+
fmt(run_model),
243+
migrated_count,
244+
)
245+
246+
179247
async def _process_pending_run(session: AsyncSession, run_model: RunModel):
180248
"""Jobs are not created yet"""
181249
run = run_model_to_run(run_model)

src/dstack/_internal/server/background/tasks/process_submitted_jobs.py

Lines changed: 34 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -740,11 +740,23 @@ def _get_profile_for_job(run_spec: RunSpec, job: Job) -> Profile:
740740
base_profile = run_spec.merged_profile
741741

742742
group_name = job.job_spec.replica_group_name
743+
logger.info(
744+
"Getting profile for job %s: replica_group_name=%s, config_type=%s",
745+
job.job_spec.job_name,
746+
group_name,
747+
run_spec.configuration.type,
748+
)
749+
743750
if not group_name or run_spec.configuration.type != "service":
751+
logger.info("Using base profile (no group_name or not a service)")
744752
return base_profile
745753

746754
# Find the group
747755
normalized_groups = get_normalized_replica_groups(run_spec.configuration)
756+
logger.info(
757+
"Normalized groups: %s",
758+
[f"{g.name} (regions={g.regions})" for g in normalized_groups],
759+
)
748760
group = next((g for g in normalized_groups if g.name == group_name), None)
749761

750762
if not group:
@@ -762,7 +774,7 @@ def _get_profile_for_job(run_spec: RunSpec, job: Job) -> Profile:
762774
if group_value is not None:
763775
setattr(merged, field_name, group_value)
764776

765-
logger.debug(
777+
logger.info(
766778
"Profile for group '%s': regions=%s, backends=%s, spot_policy=%s (base had: regions=%s, backends=%s, spot_policy=%s)",
767779
group_name,
768780
merged.regions,
@@ -832,6 +844,21 @@ async def _run_job_on_new_instance(
832844
multinode = job.job_spec.jobs_per_replica > 1 or (
833845
fleet is not None and fleet.spec.configuration.placement == InstanceGroupPlacement.CLUSTER
834846
)
847+
848+
# Log the requirements and profile being used
849+
gpu_requirement = (
850+
requirements.resources.gpu.name
851+
if requirements.resources and requirements.resources.gpu
852+
else None
853+
)
854+
logger.info(
855+
"%s: Fetching offers with GPU=%s, regions=%s, backends=%s",
856+
fmt(job_model),
857+
gpu_requirement,
858+
profile.regions,
859+
[b.value for b in profile.backends] if profile.backends else None,
860+
)
861+
835862
offers = await get_offers_by_requirements(
836863
project=project,
837864
profile=profile,
@@ -845,14 +872,12 @@ async def _run_job_on_new_instance(
845872
)
846873

847874
# Debug logging for offers
848-
if replica_group_name and len(offers) > 0:
849-
logger.debug(
850-
"%s: Got %d offers for group '%s'. First 3: %s",
851-
fmt(job_model),
852-
len(offers),
853-
replica_group_name,
854-
[f"{o.instance.name} ({o.backend.value}/{o.region})" for _, o in offers[:3]],
855-
)
875+
logger.info(
876+
"%s: Got %d offers. First 3: %s",
877+
fmt(job_model),
878+
len(offers),
879+
[f"{o.instance.name} ({o.backend.value}/{o.region})" for _, o in offers[:3]],
880+
)
856881
# Limit number of offers tried to prevent long-running processing
857882
# in case all offers fail.
858883
for backend, offer in offers[: settings.MAX_OFFERS_TRIED]:

src/dstack/_internal/server/services/runs.py

Lines changed: 50 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -417,13 +417,36 @@ async def get_plan(
417417
job_num=0,
418418
)
419419

420-
pool_offers = await _get_pool_offers(
421-
session=session,
422-
project=project,
423-
run_spec=effective_run_spec,
424-
job=jobs[0],
425-
volumes=volumes,
426-
)
420+
# For replica groups, we need pool offers for all GPU types, not just the first job's type
421+
# So we fetch pool offers for each job separately and aggregate them
422+
all_pool_offers = []
423+
if len(jobs) > 1:
424+
# Multiple jobs (likely replica groups) - get pool offers per job to include all GPU types
425+
for job in jobs:
426+
job_pool_offers = await _get_pool_offers(
427+
session=session,
428+
project=project,
429+
run_spec=effective_run_spec,
430+
job=job,
431+
volumes=volumes,
432+
)
433+
all_pool_offers.extend(job_pool_offers)
434+
# Deduplicate by (backend, instance_name, region) tuple
435+
seen_offers = set()
436+
pool_offers = []
437+
for offer in all_pool_offers:
438+
offer_key = (offer.backend, offer.instance.name, offer.region)
439+
if offer_key not in seen_offers:
440+
seen_offers.add(offer_key)
441+
pool_offers.append(offer)
442+
else:
443+
pool_offers = await _get_pool_offers(
444+
session=session,
445+
project=project,
446+
run_spec=effective_run_spec,
447+
job=jobs[0],
448+
volumes=volumes,
449+
)
427450
effective_run_spec.run_name = "dry-run" # will regenerate jobs on submission
428451

429452
# Check if all jobs have identical requirements (optimization for single-type jobs)
@@ -1561,10 +1584,29 @@ async def retry_run_replica_jobs(
15611584
session=session,
15621585
project=run_model.project,
15631586
)
1587+
1588+
# Determine which replica group this job belongs to
1589+
run_spec = RunSpec.__response__.parse_raw(run_model.run_spec)
1590+
replica_group = None
1591+
if run_spec.configuration.type == "service" and latest_jobs:
1592+
from dstack._internal.core.models.runs import get_normalized_replica_groups
1593+
1594+
group_name = latest_jobs[0].replica_group_name
1595+
if group_name:
1596+
normalized_groups = get_normalized_replica_groups(run_spec.configuration)
1597+
replica_group = next((g for g in normalized_groups if g.name == group_name), None)
1598+
if replica_group:
1599+
logger.info(
1600+
"%s: retrying job from replica group '%s'",
1601+
fmt(run_model),
1602+
replica_group.name,
1603+
)
1604+
15641605
new_jobs = await get_jobs_from_run_spec(
1565-
run_spec=RunSpec.__response__.parse_raw(run_model.run_spec),
1606+
run_spec=run_spec,
15661607
secrets=secrets,
15671608
replica_num=latest_jobs[0].replica_num,
1609+
replica_group=replica_group,
15681610
)
15691611
assert len(new_jobs) == len(latest_jobs), (
15701612
"Changing the number of jobs within a replica is not yet supported"

src/dstack/_internal/server/testing/common.py

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -354,8 +354,19 @@ async def create_job(
354354
if deployment_num is None:
355355
deployment_num = run.deployment_num
356356
run_spec = RunSpec.parse_raw(run.run_spec)
357+
358+
# Look up replica group if specified
359+
replica_group = None
360+
if replica_group_name and run_spec.configuration.type == "service":
361+
from dstack._internal.core.models.runs import get_normalized_replica_groups
362+
363+
normalized_groups = get_normalized_replica_groups(run_spec.configuration)
364+
replica_group = next((g for g in normalized_groups if g.name == replica_group_name), None)
365+
357366
job_spec = (
358-
await get_job_specs_from_run_spec(run_spec=run_spec, secrets={}, replica_num=replica_num)
367+
await get_job_specs_from_run_spec(
368+
run_spec=run_spec, secrets={}, replica_num=replica_num, replica_group=replica_group
369+
)
359370
)[0]
360371
job_spec.job_num = job_num
361372
job = JobModel(

0 commit comments

Comments
 (0)