@@ -1023,158 +1023,31 @@ def _k_idx_hi(ks):
10231023 q_start_i32 = fx .Int32 (q_start )
10241024 max_kv_col_i32 = kv_start_i32 + fx .Int32 (BLOCK_N - 1 )
10251025 tile_needs_mask = max_kv_col_i32 > q_start_i32
1026- s_raw_lo_0 = s_raw_lo [0 ]
1027- s_raw_lo_1 = s_raw_lo [1 ]
1028- s_raw_lo_2 = s_raw_lo [2 ]
1029- s_raw_lo_3 = s_raw_lo [3 ]
1030- s_raw_lo_4 = s_raw_lo [4 ]
1031- s_raw_lo_5 = s_raw_lo [5 ]
1032- s_raw_lo_6 = s_raw_lo [6 ]
1033- s_raw_lo_7 = s_raw_lo [7 ]
1034- s_raw_lo_8 = s_raw_lo [8 ]
1035- s_raw_lo_9 = s_raw_lo [9 ]
1036- s_raw_lo_10 = s_raw_lo [10 ]
1037- s_raw_lo_11 = s_raw_lo [11 ]
1038- s_raw_lo_12 = s_raw_lo [12 ]
1039- s_raw_lo_13 = s_raw_lo [13 ]
1040- s_raw_lo_14 = s_raw_lo [14 ]
1041- s_raw_lo_15 = s_raw_lo [15 ]
1042- s_raw_hi_0 = s_raw_hi [0 ]
1043- s_raw_hi_1 = s_raw_hi [1 ]
1044- s_raw_hi_2 = s_raw_hi [2 ]
1045- s_raw_hi_3 = s_raw_hi [3 ]
1046- s_raw_hi_4 = s_raw_hi [4 ]
1047- s_raw_hi_5 = s_raw_hi [5 ]
1048- s_raw_hi_6 = s_raw_hi [6 ]
1049- s_raw_hi_7 = s_raw_hi [7 ]
1050- s_raw_hi_8 = s_raw_hi [8 ]
1051- s_raw_hi_9 = s_raw_hi [9 ]
1052- s_raw_hi_10 = s_raw_hi [10 ]
1053- s_raw_hi_11 = s_raw_hi [11 ]
1054- s_raw_hi_12 = s_raw_hi [12 ]
1055- s_raw_hi_13 = s_raw_hi [13 ]
1056- s_raw_hi_14 = s_raw_hi [14 ]
1057- s_raw_hi_15 = s_raw_hi [15 ]
1058-
1026+ # NOTE: concise causal-mask form (112 hand-unrolled lines -> a loop).
1027+ # This currently FAILS at compile time:
1028+ # "state variable 's_raw_lo' is list, not an MLIR Value"
1029+ # Root cause (suspected FlyDSL defect): `tile_needs_mask` is a runtime
1030+ # value, so this is a dynamic `if` (scf.if). The if-rewriter can carry a
1031+ # reassigned *named scalar* MLIR value out of the branch, but NOT a
1032+ # reassigned *list* local. The original 112-line unroll only existed to
1033+ # work around this by carrying 16+16 named scalars. Kept here as a repro
1034+ # so the rewriter can be fixed to carry list-typed locals (flatten to
1035+ # per-element yields), after which this is the intended form.
10591036 if tile_needs_mask :
10601037 lane_off_i32 = lane_div_32_i32 * fx .Int32 (4 )
1061- kv_col_lo_0 = kv_start_i32 + lane_off_i32 + fx .Int32 (0 )
1062- s_raw_lo_0 = ArithValue (kv_col_lo_0 > q_row_i32 ).select (c_neg_inf , s_raw_lo_0 )
1063- s_raw_hi_0 = ArithValue (kv_col_lo_0 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1064- c_neg_inf , s_raw_hi_0
1065- )
1066- kv_col_lo_1 = kv_start_i32 + lane_off_i32 + fx .Int32 (1 )
1067- s_raw_lo_1 = ArithValue (kv_col_lo_1 > q_row_i32 ).select (c_neg_inf , s_raw_lo_1 )
1068- s_raw_hi_1 = ArithValue (kv_col_lo_1 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1069- c_neg_inf , s_raw_hi_1
1070- )
1071- kv_col_lo_2 = kv_start_i32 + lane_off_i32 + fx .Int32 (2 )
1072- s_raw_lo_2 = ArithValue (kv_col_lo_2 > q_row_i32 ).select (c_neg_inf , s_raw_lo_2 )
1073- s_raw_hi_2 = ArithValue (kv_col_lo_2 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1074- c_neg_inf , s_raw_hi_2
1075- )
1076- kv_col_lo_3 = kv_start_i32 + lane_off_i32 + fx .Int32 (3 )
1077- s_raw_lo_3 = ArithValue (kv_col_lo_3 > q_row_i32 ).select (c_neg_inf , s_raw_lo_3 )
1078- s_raw_hi_3 = ArithValue (kv_col_lo_3 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1079- c_neg_inf , s_raw_hi_3
1080- )
1081- kv_col_lo_4 = kv_start_i32 + lane_off_i32 + fx .Int32 (8 )
1082- s_raw_lo_4 = ArithValue (kv_col_lo_4 > q_row_i32 ).select (c_neg_inf , s_raw_lo_4 )
1083- s_raw_hi_4 = ArithValue (kv_col_lo_4 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1084- c_neg_inf , s_raw_hi_4
1085- )
1086- kv_col_lo_5 = kv_start_i32 + lane_off_i32 + fx .Int32 (9 )
1087- s_raw_lo_5 = ArithValue (kv_col_lo_5 > q_row_i32 ).select (c_neg_inf , s_raw_lo_5 )
1088- s_raw_hi_5 = ArithValue (kv_col_lo_5 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1089- c_neg_inf , s_raw_hi_5
1090- )
1091- kv_col_lo_6 = kv_start_i32 + lane_off_i32 + fx .Int32 (10 )
1092- s_raw_lo_6 = ArithValue (kv_col_lo_6 > q_row_i32 ).select (c_neg_inf , s_raw_lo_6 )
1093- s_raw_hi_6 = ArithValue (kv_col_lo_6 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1094- c_neg_inf , s_raw_hi_6
1095- )
1096- kv_col_lo_7 = kv_start_i32 + lane_off_i32 + fx .Int32 (11 )
1097- s_raw_lo_7 = ArithValue (kv_col_lo_7 > q_row_i32 ).select (c_neg_inf , s_raw_lo_7 )
1098- s_raw_hi_7 = ArithValue (kv_col_lo_7 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1099- c_neg_inf , s_raw_hi_7
1100- )
1101- kv_col_lo_8 = kv_start_i32 + lane_off_i32 + fx .Int32 (16 )
1102- s_raw_lo_8 = ArithValue (kv_col_lo_8 > q_row_i32 ).select (c_neg_inf , s_raw_lo_8 )
1103- s_raw_hi_8 = ArithValue (kv_col_lo_8 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1104- c_neg_inf , s_raw_hi_8
1105- )
1106- kv_col_lo_9 = kv_start_i32 + lane_off_i32 + fx .Int32 (17 )
1107- s_raw_lo_9 = ArithValue (kv_col_lo_9 > q_row_i32 ).select (c_neg_inf , s_raw_lo_9 )
1108- s_raw_hi_9 = ArithValue (kv_col_lo_9 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1109- c_neg_inf , s_raw_hi_9
1110- )
1111- kv_col_lo_10 = kv_start_i32 + lane_off_i32 + fx .Int32 (18 )
1112- s_raw_lo_10 = ArithValue (kv_col_lo_10 > q_row_i32 ).select (c_neg_inf , s_raw_lo_10 )
1113- s_raw_hi_10 = ArithValue (kv_col_lo_10 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1114- c_neg_inf , s_raw_hi_10
1115- )
1116- kv_col_lo_11 = kv_start_i32 + lane_off_i32 + fx .Int32 (19 )
1117- s_raw_lo_11 = ArithValue (kv_col_lo_11 > q_row_i32 ).select (c_neg_inf , s_raw_lo_11 )
1118- s_raw_hi_11 = ArithValue (kv_col_lo_11 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1119- c_neg_inf , s_raw_hi_11
1120- )
1121- kv_col_lo_12 = kv_start_i32 + lane_off_i32 + fx .Int32 (24 )
1122- s_raw_lo_12 = ArithValue (kv_col_lo_12 > q_row_i32 ).select (c_neg_inf , s_raw_lo_12 )
1123- s_raw_hi_12 = ArithValue (kv_col_lo_12 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1124- c_neg_inf , s_raw_hi_12
1125- )
1126- kv_col_lo_13 = kv_start_i32 + lane_off_i32 + fx .Int32 (25 )
1127- s_raw_lo_13 = ArithValue (kv_col_lo_13 > q_row_i32 ).select (c_neg_inf , s_raw_lo_13 )
1128- s_raw_hi_13 = ArithValue (kv_col_lo_13 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1129- c_neg_inf , s_raw_hi_13
1130- )
1131- kv_col_lo_14 = kv_start_i32 + lane_off_i32 + fx .Int32 (26 )
1132- s_raw_lo_14 = ArithValue (kv_col_lo_14 > q_row_i32 ).select (c_neg_inf , s_raw_lo_14 )
1133- s_raw_hi_14 = ArithValue (kv_col_lo_14 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1134- c_neg_inf , s_raw_hi_14
1135- )
1136- kv_col_lo_15 = kv_start_i32 + lane_off_i32 + fx .Int32 (27 )
1137- s_raw_lo_15 = ArithValue (kv_col_lo_15 > q_row_i32 ).select (c_neg_inf , s_raw_lo_15 )
1138- s_raw_hi_15 = ArithValue (kv_col_lo_15 + fx .Int32 (K_SUB_N ) > q_row_i32 ).select (
1139- c_neg_inf , s_raw_hi_15
1140- )
1141-
1142- s_raw_lo = [
1143- s_raw_lo_0 ,
1144- s_raw_lo_1 ,
1145- s_raw_lo_2 ,
1146- s_raw_lo_3 ,
1147- s_raw_lo_4 ,
1148- s_raw_lo_5 ,
1149- s_raw_lo_6 ,
1150- s_raw_lo_7 ,
1151- s_raw_lo_8 ,
1152- s_raw_lo_9 ,
1153- s_raw_lo_10 ,
1154- s_raw_lo_11 ,
1155- s_raw_lo_12 ,
1156- s_raw_lo_13 ,
1157- s_raw_lo_14 ,
1158- s_raw_lo_15 ,
1159- ]
1160- s_raw_hi = [
1161- s_raw_hi_0 ,
1162- s_raw_hi_1 ,
1163- s_raw_hi_2 ,
1164- s_raw_hi_3 ,
1165- s_raw_hi_4 ,
1166- s_raw_hi_5 ,
1167- s_raw_hi_6 ,
1168- s_raw_hi_7 ,
1169- s_raw_hi_8 ,
1170- s_raw_hi_9 ,
1171- s_raw_hi_10 ,
1172- s_raw_hi_11 ,
1173- s_raw_hi_12 ,
1174- s_raw_hi_13 ,
1175- s_raw_hi_14 ,
1176- s_raw_hi_15 ,
1177- ]
1038+ s_raw_lo = [
1039+ ArithValue (
1040+ kv_start_i32 + lane_off_i32 + fx .Int32 ((r // 4 ) * 8 + (r % 4 )) > q_row_i32
1041+ ).select (c_neg_inf , s_raw_lo [r ])
1042+ for r in range_constexpr (16 )
1043+ ]
1044+ s_raw_hi = [
1045+ ArithValue (
1046+ kv_start_i32 + lane_off_i32 + fx .Int32 ((r // 4 ) * 8 + (r % 4 )) + fx .Int32 (K_SUB_N )
1047+ > q_row_i32
1048+ ).select (c_neg_inf , s_raw_hi [r ])
1049+ for r in range_constexpr (16 )
1050+ ]
11781051 else :
11791052 # Non-causal KV padding mask: keys with absolute column >= seq_len
11801053 # -> -inf, so OOB KV (0 or duplicated row) doesn't leak into softmax.
0 commit comments