@@ -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