@@ -205,7 +205,7 @@ def group_reduce(x, mode):
205205 off = fx .Int32 (THREADS_PER_TOKEN // (2 << _sh ))
206206 peer = w .shuffle_xor (off , width_i32 )
207207 if mode == "max" :
208- w = w . maximumf ( peer )
208+ w = fx . max ( w , peer )
209209 else :
210210 w = w .addf (peer , fastmath = fm_fast )
211211 return w
@@ -314,7 +314,7 @@ def _store_scalar_i32(divided, index, val):
314314 val_e = vector .extract (as_ir_value (atom_vec ), dynamic_position = [], static_position = [v ])
315315 xv = val_e if dtype_str == "f32" else val_e .extf (compute_type )
316316 x_list .append (xv )
317- thread_max = thread_max . maximumf ( xv )
317+ thread_max = fx . max ( thread_max , xv )
318318
319319 group_max = group_reduce (thread_max , "max" )
320320
@@ -364,7 +364,7 @@ def _store_scalar_i32(divided, index, val):
364364
365365 # Pass 5: leader writes weights/indices/tei (with optional renorm).
366366 c_eps = fx .Float32 (1e-20 )
367- denom = selected_sum . maximumf ( c_eps )
367+ denom = fx . max ( selected_sum , c_eps )
368368 inv_denom = c_one_f / denom
369369
370370 if (expert_lane == fx .Int32 (0 )) & (global_token < i32_num_tokens ):
@@ -467,7 +467,7 @@ def group_reduce(x, mode):
467467 off = fx .Int32 (THREADS_PER_TOKEN // (2 << _sh ))
468468 peer = w .shuffle_xor (off , width_i32 )
469469 if mode == "max" :
470- w = w . maximumf ( peer )
470+ w = fx . max ( w , peer )
471471 else :
472472 w = w .addf (peer , fastmath = fm_fast )
473473 return w
@@ -568,7 +568,7 @@ def _store_scalar_i32(divided, index, val):
568568 val_e = vector .extract (as_ir_value (atom_vec ), dynamic_position = [], static_position = [v ])
569569 xv = val_e if dtype_str == "f32" else val_e .extf (compute_type )
570570 x_list .append (xv )
571- thread_max = thread_max . maximumf ( xv )
571+ thread_max = fx . max ( thread_max , xv )
572572
573573 group_max = group_reduce (thread_max , "max" )
574574
@@ -632,7 +632,7 @@ def _store_scalar_i32(divided, index, val):
632632 # Pass 5: Leader writes weights/indices/tei (with optional renorm)
633633 # ==================================================================
634634 c_eps = fx .Float32 (1e-20 )
635- denom = selected_sum . maximumf ( c_eps )
635+ denom = fx . max ( selected_sum , c_eps )
636636 inv_denom = c_one_f / denom
637637
638638 # Inline the leader-active predicate so the AST rewriter recognises it
0 commit comments