Skip to content

Commit e792cd7

Browse files
Weiwei Yangchanglan
authored andcommitted
Add labels to LWS configs
GitOrigin-RevId: aa81577
1 parent 691f891 commit e792cd7

2 files changed

Lines changed: 38 additions & 1 deletion

File tree

axlearn/cloud/gcp/job.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -547,12 +547,14 @@ class Config(GCPJob.Config):
547547
builder: A builder that returns one or more statefulset specs.
548548
namespace: The namespace to use within the k8s cluster.
549549
annotations: LeaderWorkerSet annotations.
550+
labels: LeaderWorkerSet labels.
550551
num_replicas: number of LWS replicas.
551552
"""
552553

553554
builder: Required[BaseLeaderWorkerTemplate.Config] = REQUIRED
554555
namespace: str = "default"
555556
annotations: Optional[ConfigOr[dict]] = None
557+
labels: Optional[ConfigOr[dict]] = None
556558
num_replicas: int = 1
557559
enable_service: bool = False
558560
ports: list[str] = None
@@ -700,7 +702,7 @@ def _build_leaderworkerset(self) -> Nested[Any]:
700702
"""
701703
cfg: GKELeaderWorkerSet.Config = self.config
702704
annotations = maybe_instantiate(cfg.annotations or {})
703-
labels = {}
705+
labels = maybe_instantiate(cfg.labels or {})
704706

705707
# If the topology is set and slice auto provisioning is configured
706708
# set the necessary annotations

axlearn/cloud/gcp/job_test.py

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -654,3 +654,38 @@ def test_build_leaderworkerset(
654654
self.assertIn("spec", lws_spec)
655655
self.assertIn("replicas", lws_spec["spec"])
656656
self.assertIn("leaderWorkerTemplate", lws_spec["spec"])
657+
658+
@parameterized.product(
659+
bundler_cls=[ArtifactRegistryBundler, CloudBuildBundler],
660+
labels=[None, {"env": "tpu-test"}, {"team": "research", "experiment": "training"}],
661+
)
662+
def test_build_leaderworkerset_labels(
663+
self,
664+
bundler_cls,
665+
labels: Optional[dict] = None,
666+
):
667+
"""Test that labels are properly set in LeaderWorkerSet metadata."""
668+
cfg, bundler_cfg = self._job_config(
669+
command="test-command",
670+
bundler_cls=bundler_cls,
671+
)
672+
# Set labels on the config
673+
cfg = cfg.set(labels=labels)
674+
gke_job: job.GKELeaderWorkerSet = cfg.instantiate(bundler=bundler_cfg.instantiate())
675+
# pylint: disable-next=protected-access
676+
lws_spec = gke_job._build_leaderworkerset()
677+
lws_metadata = lws_spec["metadata"]
678+
lws_labels = lws_metadata.get("labels", {})
679+
680+
# Test basic metadata
681+
self.assertEqual(lws_metadata["name"], cfg.name)
682+
683+
# Test labels
684+
if labels is None:
685+
# When labels is None, labels dict should be empty
686+
self.assertEqual(lws_labels, {})
687+
else:
688+
# When labels is provided, they should be present in metadata
689+
for key, value in labels.items():
690+
self.assertIn(key, lws_labels)
691+
self.assertEqual(lws_labels[key], value)

0 commit comments

Comments
 (0)