Skip to content

Commit 0aff374

Browse files
xudoyuanclaude
andcommitted
[repro] flash_attn_generic: concise causal-mask loop trips dynamic-if list-carry
Collapses the 112-line hand-unrolled causal mask (s_raw_lo_0..15 / s_raw_hi_0..15) into a range_constexpr(16) loop. This is the intended concise form but it currently FAILS at compile time: TypeError: state variable 's_raw_lo' is list, not an MLIR Value; stateful dynamic if requires MLIR-backed values. Suspected FlyDSL compile-time defect: `tile_needs_mask` is a runtime value, so `if tile_needs_mask:` lowers to a dynamic scf.if. The if-rewriter can carry a reassigned *named scalar* MLIR value out of the branch, but not a reassigned *list* local. The original 112-line unroll existed solely to work around this (carry 16+16 named scalars). Committed intentionally as a minimal repro so the rewriter can be fixed to carry list-typed locals (flatten to per-element scf.if yields). Once fixed, this is the desired form. Does NOT build as-is. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
1 parent f5c50f4 commit 0aff374

1 file changed

Lines changed: 23 additions & 150 deletions

File tree

kernels/flash_attn_generic.py

Lines changed: 23 additions & 150 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)