diff --git a/pathwaysutils/experimental/gke/jobset.py b/pathwaysutils/experimental/gke/jobset.py index 28bee85..ab2e798 100644 --- a/pathwaysutils/experimental/gke/jobset.py +++ b/pathwaysutils/experimental/gke/jobset.py @@ -124,6 +124,9 @@ def __init__( shared_pathways_service: bool = False, pathways_rm_and_worker_image: str = DEFAULT_PATHWAYS_RM_AND_WORKER_IMAGE, pathways_proxy_image: str = DEFAULT_PATHWAYS_PROXY_IMAGE, + head_nodepool: str | None = None, + head_node_selector: Mapping[str, str] | None = None, + priority_class_name: str | None = None, ): """Initializes the instance. @@ -135,17 +138,23 @@ def __init__( topology: TPU topology (e.g., "2x2"). num_slices: Number of slices. max_restarts: Maximum number of restarts for the JobSet. - max_slice_restarts: Maximum number of slice restarts (defaults to 1_000_000 in headless and SPS mode). + max_slice_restarts: Maximum number of slice restarts (defaults to + 1_000_000 in headless and SPS mode). termination_grace_period_seconds: Optional termination grace period. pathways_version: Version tag for Pathways images. jobset_api_version: API version of JobSet. elastic_slices: Number of elastic slices. labels: Optional labels for the JobSet. annotations: Optional annotations for the JobSet. - shared_pathways_service: Whether to run only RM for Shared Pathways Service. + shared_pathways_service: Whether to run only RM for Shared Pathways + Service. pathways_rm_and_worker_image: Base Docker image for Resource Manager and Worker containers. pathways_proxy_image: Base Docker image for Proxy container. + head_nodepool: Optional GKE nodepool for the head pod. + head_node_selector: Optional node selector dict for the head pod. + priority_class_name: Optional priority class name for head and worker + pods. """ self._shared_pathways_service = shared_pathways_service self._pathways_rm_and_worker_image = pathways_rm_and_worker_image @@ -194,6 +203,9 @@ def __init__( shared_pathways_service=shared_pathways_service, pathways_rm_and_worker_image=pathways_rm_and_worker_image, pathways_proxy_image=pathways_proxy_image, + head_nodepool=head_nodepool, + head_node_selector=head_node_selector, + priority_class_name=priority_class_name, ) # Build worker template. @@ -207,14 +219,10 @@ def __init__( max_slice_restarts=max_slice_restarts, termination_grace_period_seconds=termination_grace_period_seconds, pathways_rm_and_worker_image=pathways_rm_and_worker_image, + priority_class_name=priority_class_name, ) self._success_policy = None - if shared_pathways_service: - self._success_policy = { - "operator": "All", - "targetReplicatedJobs": [PATHWAYS_HEAD_JOB_NAME], - } @property def head_job_template(self) -> client.V1JobTemplateSpec: @@ -242,6 +250,9 @@ def _build_head_job_template( shared_pathways_service: bool, pathways_rm_and_worker_image: str, pathways_proxy_image: str, + head_nodepool: str | None = None, + head_node_selector: Mapping[str, str] | None = None, + priority_class_name: str | None = None, ) -> client.V1JobTemplateSpec: """Builds the head job template for the JobSet. @@ -251,9 +262,13 @@ def _build_head_job_template( instance_type: TPU instance type (e.g., "tpuv5:2x2"). image_tag: Version tag for Pathways images. elastic_slices: Number of elastic slices. - shared_pathways_service: Whether to run only RM for Shared Pathways Service. + shared_pathways_service: Whether to run only RM for Shared Pathways + Service. pathways_rm_and_worker_image: Base Docker image for Resource Manager. pathways_proxy_image: Base Docker image for Proxy container. + head_nodepool: Optional GKE nodepool for the head pod. + head_node_selector: Optional node selector dict for the head pod. + priority_class_name: Optional priority class name. Returns: The head job template. @@ -352,14 +367,34 @@ def _build_head_job_template( ), ) - containers = [rm_container] + rm_container.restart_policy = "Always" + proxy_container.restart_policy = "Always" + + init_containers = [rm_container] if not shared_pathways_service: - containers.append(proxy_container) + init_containers.append(proxy_container) + + dummy_container = client.V1Container( + name="dummy-job", + image="gcr.io/google-containers/pause:3.2", + ) head_pod_spec = client.V1PodSpec( - containers=containers, + init_containers=init_containers, + containers=[dummy_container], restart_policy="Never", + dns_policy="ClusterFirstWithHostNet", + host_network=True, ) + node_sel = dict(head_node_selector) if head_node_selector else {} + if head_nodepool: + node_sel["cloud.google.com/gke-nodepool"] = head_nodepool + elif "cloud.google.com/gke-nodepool" not in node_sel: + node_sel["cloud.google.com/gke-nodepool"] = "cpu-np" + head_pod_spec.node_selector = node_sel + + if priority_class_name: + head_pod_spec.priority_class_name = priority_class_name job_annotations = { "alpha.jobset.sigs.k8s.io/exclusive-topology": "kubernetes.io/hostname" @@ -393,6 +428,7 @@ def _build_worker_job_template( max_slice_restarts: int, termination_grace_period_seconds: int | None, pathways_rm_and_worker_image: str, + priority_class_name: str | None = None, ) -> client.V1JobTemplateSpec: """Builds the worker job template for the JobSet.""" worker_image = _format_image(pathways_rm_and_worker_image, image_tag) @@ -504,7 +540,11 @@ def _build_worker_job_template( ) ], restart_policy="OnFailure", + dns_policy="ClusterFirstWithHostNet", + host_network=True, ) + if priority_class_name: + worker_pod_spec.priority_class_name = priority_class_name if termination_grace_period_seconds is not None: worker_pod_spec.termination_grace_period_seconds = ( termination_grace_period_seconds @@ -574,6 +614,93 @@ def _add_volume_to_pod_spec( volumes.append(volume) pod_spec.volumes = volumes + def add_user_workload( + self, + image: str, + command: Sequence[str] | str, + name: str = "user-workload", + env: Mapping[str, str] | None = None, + resources: client.V1ResourceRequirements | None = None, + ) -> "PathwaysJobSet": + """Adds a user workload container to the head pod. + + Args: + image: Docker image for the user workload. + command: Command to execute in the container (string or sequence of + strings). + name: Container name for the user workload (default: 'user-workload'). + env: Optional additional environment variables. + resources: Optional resource requirements (default: 24 CPUs, 100G memory). + + Returns: + The PathwaysJobSet instance for chaining. + """ + pod_spec = self._head_job_template.spec.template.spec + + cmd = ["sh", "-c", command] if isinstance(command, str) else list(command) + + user_env_list = [ + client.V1EnvVar(name="JAX_PLATFORMS", value="proxy"), + client.V1EnvVar( + name="PATHWAYS_HEAD", + value_from=client.V1EnvVarSource( + field_ref=client.V1ObjectFieldSelector( + field_path=( + "metadata.labels['jobset.sigs.k8s.io/coordinator']" + ) + ) + ), + ), + client.V1EnvVar( + name="JAX_BACKEND_TARGET", + value=f"grpc://localhost:{PATHWAYS_PROXY_PORT}", + ), + client.V1EnvVar( + name="MEGASCALE_NUM_SLICES", + value=str(self._worker_replicas), + ), + client.V1EnvVar( + name="JOBSET_NAME", + value_from=client.V1EnvVarSource( + field_ref=client.V1ObjectFieldSelector( + field_path=( + "metadata.annotations['jobset.sigs.k8s.io/jobset-name']" + ) + ) + ), + ), + ] + if env: + for k, v in env.items(): + user_env_list.append(client.V1EnvVar(name=k, value=str(v))) + + if resources is None: + resources = client.V1ResourceRequirements( + limits={"cpu": "24", "memory": "100G"} + ) + + user_container = client.V1Container( + name=name, + image=image, + image_pull_policy="Always", + command=cmd, + env=user_env_list, + resources=resources, + ) + + containers = pod_spec.containers or [] + containers = [c for c in containers if c.name != name and c.name != "dummy-job"] + containers.append(user_container) + pod_spec.containers = containers + + if not self._shared_pathways_service: + self._success_policy = { + "operator": "All", + "targetReplicatedJobs": [PATHWAYS_HEAD_JOB_NAME], + } + + return self + def add_colocated_python( self, image: str, diff --git a/pathwaysutils/test/experimental/gke/jobset_test.py b/pathwaysutils/test/experimental/gke/jobset_test.py index 41fd382..3a21b01 100644 --- a/pathwaysutils/test/experimental/gke/jobset_test.py +++ b/pathwaysutils/test/experimental/gke/jobset_test.py @@ -149,8 +149,8 @@ def test_headless_head_job_pod_spec(self): self.assertIn("pathways-head", helper.jobs) self.assertEqual(helper.jobs["pathways-head"]["replicas"], 1) pod_spec = helper.pod_specs["pathways-head"] - self.assertNotIn("hostNetwork", pod_spec) - self.assertNotIn("dnsPolicy", pod_spec) + self.assertTrue(pod_spec["hostNetwork"]) + self.assertEqual(pod_spec["dnsPolicy"], "ClusterFirstWithHostNet") self.assertEqual(pod_spec["restartPolicy"], "Never") def test_headless_head_job_containers(self): @@ -160,17 +160,19 @@ def test_headless_head_job_containers(self): helper = JobSetManifestHelper(config) pod_spec = helper.pod_specs["pathways-head"] - self.assertLen(pod_spec["containers"], 2) - self.assertIn("pathways-rm", helper.containers["pathways-head"]) - self.assertIn("pathways-proxy", helper.containers["pathways-head"]) + self.assertLen(pod_spec["initContainers"], 2) + self.assertLen(pod_spec["containers"], 1) + self.assertIn("pathways-rm", helper.init_containers["pathways-head"]) + self.assertIn("pathways-proxy", helper.init_containers["pathways-head"]) + self.assertIn("dummy-job", helper.containers["pathways-head"]) - rm_container = helper.containers["pathways-head"]["pathways-rm"] + rm_container = helper.init_containers["pathways-head"]["pathways-rm"] self.assertEqual( rm_container["image"], "us-docker.pkg.dev/cloud-tpu-v2-images/pathways/server:latest", ) - proxy_container = helper.containers["pathways-head"]["pathways-proxy"] + proxy_container = helper.init_containers["pathways-head"]["pathways-proxy"] self.assertEqual( proxy_container["image"], "us-docker.pkg.dev/cloud-tpu-v2-images/pathways/proxy_server:latest", @@ -207,8 +209,8 @@ def test_worker_job_pod_spec(self): helper = JobSetManifestHelper(config) pod_spec = helper.pod_specs["pathways-worker"] - self.assertNotIn("hostNetwork", pod_spec) - self.assertNotIn("dnsPolicy", pod_spec) + self.assertTrue(pod_spec["hostNetwork"]) + self.assertEqual(pod_spec["dnsPolicy"], "ClusterFirstWithHostNet") self.assertEqual(pod_spec["restartPolicy"], "OnFailure") self.assertEqual(pod_spec["terminationGracePeriodSeconds"], 60) @@ -266,7 +268,7 @@ def test_add_gcsfuse_read_only(self, read_only): ) helper = JobSetManifestHelper(pw_jobset.to_dict()) - rm_container = helper.containers["pathways-head"]["pathways-rm"] + rm_container = helper.init_containers["pathways-head"]["pathways-rm"] mounts = {m["name"]: m for m in rm_container.get("volumeMounts", [])} self.assertIn(expected_vol_name, mounts) self.assertEqual(mounts[expected_vol_name].get("readOnly", False), read_only) @@ -677,7 +679,7 @@ def test_configurable_images(self): # Verify head job containers. head_containers = { c.name: c.image - for c in pw_jobset.head_job_template.spec.template.spec.containers + for c in pw_jobset.head_job_template.spec.template.spec.init_containers } self.assertEqual(head_containers["pathways-rm"], "gcr.io/custom/server:v2.0") self.assertEqual( @@ -826,6 +828,32 @@ def test_shared_pathways_service(self): config = pw_jobset.to_dict() helper = JobSetManifestHelper(config) + self.assertNotIn("successPolicy", config["spec"]) + + self.assertTrue(config["spec"]["network"]["enableDNSHostnames"]) + self.assertTrue(config["spec"]["network"]["publishNotReadyAddresses"]) + + self.assertIn("pathways-head", helper.jobs) + pod_spec = helper.pod_specs["pathways-head"] + + # Head job should have pathways-rm in init_containers, no pathways-proxy, and dummy-job in containers. + self.assertIn("pathways-rm", helper.init_containers["pathways-head"]) + self.assertNotIn("pathways-proxy", helper.init_containers["pathways-head"]) + self.assertIn("dummy-job", helper.containers["pathways-head"]) + self.assertLen(pod_spec["containers"], 1) + + def test_add_user_workload(self): + pw_jobset = self._create_jobset(topology="2x2", num_slices=1) + pw_jobset.add_user_workload( + image="us-docker.pkg.dev/my-project/test:v1", + command="python3 -m test_module", + env={"CUSTOM_ENV": "123"}, + ) + + config = pw_jobset.to_dict() + helper = JobSetManifestHelper(config) + + # Verify JobSet spec-level success policy targets pathways-head self.assertEqual( config["spec"]["successPolicy"], { @@ -834,16 +862,159 @@ def test_shared_pathways_service(self): }, ) - self.assertTrue(config["spec"]["network"]["enableDNSHostnames"]) - self.assertTrue(config["spec"]["network"]["publishNotReadyAddresses"]) + # Verify head pod has user-workload container in containers + self.assertIn("user-workload", helper.containers["pathways-head"]) + user_c = helper.containers["pathways-head"]["user-workload"] + self.assertEqual(user_c["image"], "us-docker.pkg.dev/my-project/test:v1") + self.assertEqual(user_c["command"], ["sh", "-c", "python3 -m test_module"]) + self.assertTrue( + any( + e["name"] == "CUSTOM_ENV" and e["value"] == "123" + for e in user_c["env"] + ) + ) + self.assertTrue( + any( + e["name"] == "MEGASCALE_NUM_SLICES" and e["value"] == "1" + for e in user_c["env"] + ) + ) + self.assertTrue( + any( + e["name"] == "JAX_PLATFORMS" and e["value"] == "proxy" + for e in user_c["env"] + ) + ) + self.assertTrue( + any( + e["name"] == "PATHWAYS_HEAD" + and e.get("valueFrom", {}) + .get("fieldRef", {}) + .get("fieldPath") + == "metadata.labels['jobset.sigs.k8s.io/coordinator']" + for e in user_c["env"] + ) + ) + self.assertTrue( + any( + e["name"] == "JAX_BACKEND_TARGET" + and e["value"] == "grpc://localhost:29000" + for e in user_c["env"] + ) + ) - self.assertIn("pathways-head", helper.jobs) - pod_spec = helper.pod_specs["pathways-head"] + # Verify RM and Proxy are in initContainers (native sidecars) + self.assertIn("pathways-rm", helper.init_containers["pathways-head"]) + self.assertIn("pathways-proxy", helper.init_containers["pathways-head"]) + self.assertIn("user-workload", helper.containers["pathways-head"]) + self.assertNotIn("dummy-job", helper.containers["pathways-head"]) - # Head job should only have pathways-rm container, no pathways-proxy. - self.assertIn("pathways-rm", helper.containers["pathways-head"]) - self.assertNotIn("pathways-proxy", helper.containers["pathways-head"]) - self.assertLen(pod_spec["containers"], 1) + def test_add_user_workload_roundtrip(self): + pw_jobset = self._create_jobset(topology="2x2", num_slices=1) + pw_jobset.add_user_workload( + image="us-docker.pkg.dev/my-project/test:v1", + command=["python3", "test.py"], + ) + + temp_filepath = os.path.join( + self.create_tempdir().full_path, "jobset_user_workload.yaml" + ) + pw_jobset.export_yaml(temp_filepath) + imported = jobset.PathwaysJobSet.import_yaml(temp_filepath) + + self.assertEqual( + normalize_k8s_spec(pw_jobset.to_dict()), + normalize_k8s_spec(imported.to_dict()), + ) + + def test_configurable_priority_class_and_head_node_selector(self): + # Default: no priority class or head pod node selector + default_js = self._create_jobset() + default_helper = JobSetManifestHelper(default_js.to_dict()) + self.assertNotIn( + "priorityClassName", default_helper.pod_specs["pathways-head"] + ) + self.assertEqual( + default_helper.pod_specs["pathways-head"]["nodeSelector"]["cloud.google.com/gke-nodepool"], + "cpu-np", + ) + self.assertNotIn( + "priorityClassName", default_helper.pod_specs["pathways-worker"] + ) + + # Configured: priority class and head node selector set + configured_js = self._create_jobset( + head_nodepool="cpu-np", + head_node_selector={"zone": "us-central1-a"}, + priority_class_name="high", + ) + configured_helper = JobSetManifestHelper(configured_js.to_dict()) + head_spec = configured_helper.pod_specs["pathways-head"] + worker_spec = configured_helper.pod_specs["pathways-worker"] + + self.assertEqual(head_spec["priorityClassName"], "high") + self.assertEqual(worker_spec["priorityClassName"], "high") + self.assertEqual( + head_spec["nodeSelector"]["cloud.google.com/gke-nodepool"], "cpu-np" + ) + self.assertEqual(head_spec["nodeSelector"]["zone"], "us-central1-a") + + def test_kokoro_pretraining_workload_generation(self): + """Verifies that PathwaysJobSet can generate a JobSet matching Kokoro pretraining test workloads.""" + pw_jobset = jobset.PathwaysJobSet( + name="maxtext-pretraining-test", + namespace="default", + pathways_dir="gs://my-bucket/scratch", + tpu_type="v5e", + topology="4x8", + num_slices=1, + labels={"kueue.x-k8s.io/queue-name": "multislice-queue"}, + head_nodepool="cpu-np", + priority_class_name="high", + ) + pw_jobset.add_user_workload( + image="us-docker.pkg.dev/my-project/maxtext:latest", + command="python3 MaxText/train.py MaxText/configs/base.yml", + ) + + config = pw_jobset.to_dict() + helper = JobSetManifestHelper(config) + + # Verify queue label + self.assertEqual( + config["metadata"]["labels"]["kueue.x-k8s.io/queue-name"], + "multislice-queue", + ) + + # Verify success policy + self.assertEqual( + config["spec"]["successPolicy"]["targetReplicatedJobs"], + ["pathways-head"], + ) + + # Verify head pod spec: priorityClassName, nodeSelector + head_pod_spec = helper.pod_specs["pathways-head"] + self.assertEqual(head_pod_spec["priorityClassName"], "high") + self.assertEqual( + head_pod_spec["nodeSelector"]["cloud.google.com/gke-nodepool"], + "cpu-np", + ) + + # Verify head containers: RM and Proxy are in initContainers, user workload in containers + self.assertIn("pathways-rm", helper.init_containers["pathways-head"]) + self.assertIn("pathways-proxy", helper.init_containers["pathways-head"]) + self.assertIn("user-workload", helper.containers["pathways-head"]) + self.assertNotIn("dummy-job", helper.containers["pathways-head"]) + + user_container = helper.containers["pathways-head"]["user-workload"] + self.assertEqual( + user_container["command"], + ["sh", "-c", "python3 MaxText/train.py MaxText/configs/base.yml"], + ) + + # Verify worker pod spec: priorityClassName + worker_pod_spec = helper.pod_specs["pathways-worker"] + self.assertEqual(worker_pod_spec["priorityClassName"], "high") if __name__ == "__main__":