Skip to content

Commit 847a0c3

Browse files
superbobryGoogle-ML-Automation
authored andcommitted
[pallas:mosaic] Avoid passing the deprecated device_id_type= in :tpu_pallas_interpret_distributed_test
PiperOrigin-RevId: 991040247
1 parent 3243e2b commit 847a0c3

1 file changed

Lines changed: 16 additions & 54 deletions

File tree

‎tests/pallas/tpu_pallas_interpret_distributed_test.py‎

Lines changed: 16 additions & 54 deletions
Original file line numberDiff line numberDiff line change
@@ -57,12 +57,6 @@ def setUp(self):
5757
message='jax.experimental.pallas.DeviceIdType is deprecated',
5858
)
5959
)
60-
self.enter_context(
61-
jtu.ignore_warning(
62-
category=DeprecationWarning,
63-
message='device_id_type is deprecated',
64-
)
65-
)
6660

6761
if not jtu.test_device_matches(['cpu']):
6862
self.skipTest('CPU-only test')
@@ -91,14 +85,8 @@ def right_permute_kernel(input_ref, output_ref, send_sem, recv_sem):
9185
right_neighbor = lax.rem(my_id + 1, jnp.int32(num_devices))
9286

9387
barrier_sem = pltpu.get_barrier_semaphore()
94-
pl.semaphore_signal(
95-
barrier_sem,
96-
device_id=(left_neighbor,),
97-
device_id_type=pl.DeviceIdType.MESH)
98-
pl.semaphore_signal(
99-
barrier_sem,
100-
device_id=(right_neighbor,),
101-
device_id_type=pl.DeviceIdType.MESH)
88+
pl.semaphore_signal(barrier_sem, device_id=(left_neighbor,))
89+
pl.semaphore_signal(barrier_sem, device_id=(right_neighbor,))
10290
pl.semaphore_wait(barrier_sem, 2)
10391

10492
remote_copy_op = pltpu.make_async_remote_copy(
@@ -107,7 +95,6 @@ def right_permute_kernel(input_ref, output_ref, send_sem, recv_sem):
10795
send_sem=send_sem,
10896
recv_sem=recv_sem,
10997
device_id=(right_neighbor,),
110-
device_id_type=pl.DeviceIdType.MESH,
11198
)
11299
remote_copy_op.start()
113100
remote_copy_op.wait()
@@ -192,13 +179,11 @@ def _():
192179
barrier_sem,
193180
inc=1,
194181
device_id=(left_neighbor,),
195-
device_id_type=pl.DeviceIdType.MESH,
196182
)
197183
pl.semaphore_signal(
198184
barrier_sem,
199185
inc=1,
200186
device_id=(right_neighbor,),
201-
device_id_type=pl.DeviceIdType.MESH,
202187
)
203188
pl.semaphore_wait(barrier_sem, 2)
204189

@@ -220,8 +205,7 @@ def _():
220205
send_sem=send_sem,
221206
recv_sem=recv_sems.at[outer_step],
222207
device_id=(right_neighbor,),
223-
device_id_type=pl.DeviceIdType.MESH,
224-
)
208+
)
225209
remote_copy_op.start()
226210
remote_copy_op.wait()
227211

@@ -319,13 +303,11 @@ def _():
319303
barrier_sem,
320304
inc=1,
321305
device_id=(left_neighbor,),
322-
device_id_type=pl.DeviceIdType.MESH,
323306
)
324307
pl.semaphore_signal(
325308
barrier_sem,
326309
inc=1,
327310
device_id=(right_neighbor,),
328-
device_id_type=pl.DeviceIdType.MESH,
329311
)
330312
pl.semaphore_wait(barrier_sem, 2)
331313

@@ -338,7 +320,6 @@ def _():
338320
send_sem=remote_send_sem,
339321
recv_sem=remote_recv_sem,
340322
device_id=(right_neighbor,),
341-
device_id_type=pl.DeviceIdType.MESH,
342323
)
343324
initial_copy.start()
344325
initial_copy.wait()
@@ -350,8 +331,7 @@ def _():
350331
capacity_sem,
351332
inc=1,
352333
device_id=(left_neighbor,),
353-
device_id_type=pl.DeviceIdType.MESH,
354-
)
334+
)
355335

356336
# Copy the partial result our left neighbor sent to us into VMEM for
357337
# computation.
@@ -371,8 +351,7 @@ def _():
371351
send_sem=remote_send_sem,
372352
recv_sem=remote_recv_sem,
373353
device_id=(right_neighbor,),
374-
device_id_type=pl.DeviceIdType.MESH,
375-
)
354+
)
376355
remote_copy.start()
377356
# Finish local copy and accumulate while remote_copy is happening.
378357
local_copy.wait()
@@ -476,7 +455,6 @@ def signal(left_or_right, semaphore):
476455
semaphore,
477456
inc=1,
478457
device_id=(neighbor,),
479-
device_id_type=pl.DeviceIdType.MESH,
480458
)
481459

482460
def reduce_scatter_kernel(
@@ -519,13 +497,11 @@ def _():
519497
barrier_sem,
520498
inc=1,
521499
device_id=(left_neighbor,),
522-
device_id_type=pl.DeviceIdType.MESH,
523500
)
524501
pl.semaphore_signal(
525502
barrier_sem,
526503
inc=1,
527504
device_id=(right_neighbor,),
528-
device_id_type=pl.DeviceIdType.MESH,
529505
)
530506
pl.semaphore_wait(barrier_sem, 2)
531507

@@ -535,26 +511,23 @@ def _():
535511
send_sem=left_send_sem,
536512
recv_sem=left_recv_sem,
537513
device_id=(left_neighbor,),
538-
device_id_type=pl.DeviceIdType.MESH,
539-
)
514+
)
540515

541516
initial_right_copy = pltpu.make_async_remote_copy(
542517
src_ref=x_ref.at[my_id, right_copy_slice],
543518
dst_ref=hbm_scratch.at[working_slot, right_copy_slice],
544519
send_sem=right_send_sem,
545520
recv_sem=right_recv_sem,
546521
device_id=(right_neighbor,),
547-
device_id_type=pl.DeviceIdType.MESH,
548-
)
522+
)
549523

550524
left_copy = pltpu.make_async_remote_copy(
551525
src_ref=hbm_scratch.at[working_slot, left_copy_slice],
552526
dst_ref=hbm_scratch.at[receiving_slot, left_copy_slice],
553527
send_sem=left_send_sem,
554528
recv_sem=left_recv_sem,
555529
device_id=(left_neighbor,),
556-
device_id_type=pl.DeviceIdType.MESH,
557-
)
530+
)
558531
right_copy = pltpu.make_async_remote_copy(
559532
# Note: Right copy is flipped with regards to slots since we are copying
560533
# to the next outer_step iteration.
@@ -563,8 +536,7 @@ def _():
563536
send_sem=right_send_sem,
564537
recv_sem=right_recv_sem,
565538
device_id=(right_neighbor,),
566-
device_id_type=pl.DeviceIdType.MESH,
567-
)
539+
)
568540

569541
# --- Prologue ---
570542
@pl.when(is_start)
@@ -753,7 +725,6 @@ def local_barrier(left_neighbor, right_neighbor, double_barrier=True):
753725
barrier_sem,
754726
inc=1,
755727
device_id=(neighbor,),
756-
device_id_type=pl.DeviceIdType.MESH,
757728
)
758729
pl.semaphore_wait(barrier_sem, 2)
759730
if double_barrier:
@@ -771,8 +742,7 @@ def _(second_barrier):
771742
second_barrier,
772743
inc=1,
773744
device_id=(neighbor,),
774-
device_id_type=pl.DeviceIdType.MESH,
775-
)
745+
)
776746
pl.semaphore_wait(second_barrier, 2)
777747

778748
# We pick a large outer kernel block size that we do not want to place
@@ -817,7 +787,6 @@ def signal(left_or_right, semaphore):
817787
semaphore,
818788
inc=1,
819789
device_id=(neighbor,),
820-
device_id_type=pl.DeviceIdType.MESH,
821790
)
822791

823792
def reduce_scatter_kernel(
@@ -857,34 +826,30 @@ def reduce_scatter_kernel(
857826
send_sem=left_send_sem,
858827
recv_sem=left_recv_sem,
859828
device_id=(left_neighbor,),
860-
device_id_type=pl.DeviceIdType.MESH,
861-
)
829+
)
862830

863831
initial_right_copy = pltpu.make_async_remote_copy(
864832
src_ref=x_ref.at[my_id, right_copy_slice],
865833
dst_ref=hbm_scratch.at[working_slot, right_copy_slice],
866834
send_sem=right_send_sem,
867835
recv_sem=right_recv_sem,
868836
device_id=(right_neighbor,),
869-
device_id_type=pl.DeviceIdType.MESH,
870-
)
837+
)
871838

872839
left_copy = pltpu.make_async_remote_copy(
873840
src_ref=hbm_scratch.at[working_slot, left_copy_slice],
874841
dst_ref=hbm_scratch.at[receiving_slot, left_copy_slice],
875842
send_sem=left_send_sem,
876843
recv_sem=left_recv_sem,
877844
device_id=(left_neighbor,),
878-
device_id_type=pl.DeviceIdType.MESH,
879-
)
845+
)
880846
right_copy = pltpu.make_async_remote_copy(
881847
src_ref=hbm_scratch.at[receiving_slot, right_copy_slice],
882848
dst_ref=hbm_scratch.at[working_slot, right_copy_slice],
883849
send_sem=right_send_sem,
884850
recv_sem=right_recv_sem,
885851
device_id=(right_neighbor,),
886-
device_id_type=pl.DeviceIdType.MESH,
887-
)
852+
)
888853

889854
# --- Prologue ---
890855
@pl.when(is_start)
@@ -1084,8 +1049,7 @@ def _():
10841049
send_sem=send_sem,
10851050
recv_sem=recv_sem,
10861051
device_id=(dst_id,),
1087-
device_id_type=pl.DeviceIdType.MESH,
1088-
)
1052+
)
10891053
dma.start()
10901054
dma.wait_send()
10911055
recv_count += jnp.where(dst_id == my_id, 1, 0)
@@ -1099,8 +1063,7 @@ def _():
10991063
send_sem=send_sem,
11001064
recv_sem=recv_sem,
11011065
device_id=(my_id,),
1102-
device_id_type=pl.DeviceIdType.MESH,
1103-
)
1066+
)
11041067
fake_dma.wait_recv()
11051068

11061069
@jax.jit
@@ -1221,7 +1184,6 @@ def _():
12211184
dma_sems.at[0],
12221185
dma_sems.at[1],
12231186
device_id=(left_neighbor,),
1224-
device_id_type=pl.DeviceIdType.MESH,
12251187
).wait()
12261188

12271189
run = shard_map.shard_map(

0 commit comments

Comments
 (0)