@@ -127,12 +127,14 @@ def run_job(
127127 project_ssh_private_key : str ,
128128 volumes : list [Volume ],
129129 placement_group : Optional [PlacementGroup ],
130+ requirements : Requirements ,
130131 ) -> JobProvisioningData :
131132 compute_provisioning_data = self ._run_slurm_job (
132133 run = run ,
133134 job = job ,
134135 instance_offer = instance_offer ,
135136 project_ssh_public_key = project_ssh_public_key ,
137+ requirements = requirements ,
136138 )
137139 return compute_provisioning_data .job_provisioning_datas [0 ]
138140
@@ -144,13 +146,15 @@ def run_jobs(
144146 project_ssh_public_key : str ,
145147 project_ssh_private_key : str ,
146148 placement_group : Optional [PlacementGroup ],
149+ requirements : Requirements ,
147150 ) -> ComputeGroupProvisioningData :
148151 master_job = job_configurations [0 ].job
149152 return self ._run_slurm_job (
150153 run = run ,
151154 job = master_job ,
152155 instance_offer = instance_offer ,
153156 project_ssh_public_key = project_ssh_public_key ,
157+ requirements = requirements ,
154158 )
155159
156160 def terminate_instance (
@@ -172,6 +176,7 @@ def _run_slurm_job(
172176 job : Job ,
173177 instance_offer : InstanceOfferWithAvailability ,
174178 project_ssh_public_key : str ,
179+ requirements : Requirements ,
175180 ) -> ComputeGroupProvisioningData :
176181 if job .job_spec .registry_auth is not None :
177182 self ._skip_offer_cache .add (run , job , instance_offer )
@@ -196,7 +201,7 @@ def _run_slurm_job(
196201 authorized_keys = [project_ssh_public_key .strip (), run .run_spec .ssh_key_pub .strip ()]
197202
198203 node_count = job .job_spec .jobs_per_replica
199- resources_spec = job . job_spec . requirements .resources
204+ resources_spec = requirements .resources
200205 requested_resources = get_requested_resources_from_resources_spec (resources_spec )
201206
202207 partitions = _get_cluster_partitions (cluster , requested_resources )
0 commit comments