diff --git a/spada/syntax/csl/statements.py b/spada/syntax/csl/statements.py index 31761b9a..5b211199 100644 --- a/spada/syntax/csl/statements.py +++ b/spada/syntax/csl/statements.py @@ -250,6 +250,229 @@ def emit_assignment(statement: spir.AssignmentStatement, dsds: UniqueDSDDict, dt return dsd_ops.DSD_ASSIGNMENT_MAPPING[dsd_op]() +# [P4] Targeted vectorization of the nested multiply-accumulate loop +# for (k, l) in [0:Kk, 0:Kl]: Z[k] = Z[k] + A[k*Ck + l*Cl + C0] * X[f(l)] +# (declared in either variable order) into the handwritten CSL idiom (strided base DSD + per-l +# @increment_dsd_offset + one of @fmach/@fmachs/@fmacs, matching FMADSDOp's own f16/f32 dtype +# dispatch). Conservative: fires only on this exact 2-level MAC shape with Z distinct from A and +# X; anything else falls through to scalar loops. @increment_dsd_offset is available from SDK 1.x +# (used by Cerebras' own csl-examples v1.4.0 cholesky benchmark) and its elem_type parameter is +# documented to accept f16 as well as f32 (sdk.cerebras.ai/csl/language/dsds). Generalized +# ND-loop / any-DSD-op vectorization (issue #69 asks 1 and 4) is out of scope here -- only the +# loop-variable order (ask 2), dtype restriction (ask 3), and the --disable-dsd fold-in (ask 5) +# are addressed; this pass is gated solely by dsd_ops.DISABLE_DSD, with no separate flag. +_VEC_COUNTER = [0] + + +def _affine_of(expr, varnames): + """expr 를 sum(coeff[var]*var) + const 로 분해. 실패 시 None. 계수·상수는 int.""" + e = expr.value if isinstance(expr, spir.Expression) else expr + if isinstance(e, spir.Identifier): + if e.name in varnames: + return {e.name: 1}, 0 + return None + if isinstance(e, (spir.ConstantLiteral, spir.Parameter)): + try: + v = e.eval() + except Exception: + return None + return ({}, int(v)) if isinstance(v, int) else None + if isinstance(e, spir.UnaryOperator) and getattr(e, 'op', None) == '-': + sub = _affine_of(e.operand, varnames) if hasattr(e, 'operand') else None + if sub is None: + return None + return {k: -c for k, c in sub[0].items()}, -sub[1] + if isinstance(e, spir.BinaryOperator): + left = _affine_of(e.left, varnames) + right = _affine_of(e.right, varnames) + if e.op == '+' and left and right: + coeffs = dict(left[0]) + for k, c in right[0].items(): + coeffs[k] = coeffs.get(k, 0) + c + return coeffs, left[1] + right[1] + if e.op == '-' and left and right: + coeffs = dict(left[0]) + for k, c in right[0].items(): + coeffs[k] = coeffs.get(k, 0) - c + return coeffs, left[1] - right[1] + if e.op == '*' and left and right: + if not left[0]: # 왼쪽이 순수 상수 + s = left[1] + return {k: c * s for k, c in right[0].items()}, right[1] * s + if not right[0]: # 오른쪽이 순수 상수 + s = right[1] + return {k: c * s for k, c in left[0].items()}, left[1] * s + return None + return None + + +def _ids_in(expr): + e = expr.value if isinstance(expr, spir.Expression) else expr + return {n.name for n in e.walk() if isinstance(n, spir.Identifier)} + + +def _single_index(node): + """ArraySlice 가 1차원 단일 인덱스 접근이면 그 인덱스 Expression, 아니면 None.""" + if not isinstance(node, spir.ArraySlice) or len(node.indices) != 1: + return None + idx = node.indices[0] + if isinstance(idx, spir.RangeExpression): + return None + return idx + + +def _mac_vectorization_op(mat_vec_dtype: spir.ScalarType, dst_dtype: spir.ScalarType) -> Optional[str]: + """ + Selects the @fmac* builtin for the vectorized MAC idiom, mirroring FMADSDOp._as_csl's dtype + dispatch (dsd_ops.py) exactly: the two multiplicands (mat, vec) must share a dtype + (``mat_vec_dtype``, checked by the caller), paired against the accumulator dtype + (``dst_dtype`` -- here always equal to the destination dtype, since Z[k] is both the + accumulator and the destination in this idiom). + + No integer MAC builtin is present in DSD_ASSIGNMENT_MAPPING -- only @fmach, @fmachs and + @fmacs map to FMADSDOp there -- so integer dtypes intentionally return None: MAC + vectorization stays float-only (f16/f32), matching what FMADSDOp itself supports. + """ + if mat_vec_dtype == spir.ScalarType.f16 and dst_dtype == spir.ScalarType.f16: + return '@fmach' + if mat_vec_dtype == spir.ScalarType.f32 and dst_dtype == spir.ScalarType.f16: + return '@fmachs' + if mat_vec_dtype == spir.ScalarType.f32 and dst_dtype == spir.ScalarType.f32: + return '@fmacs' + return None + + +def _try_emit_vectorized_mac(statement: spir.ForStatement, dtypes: dict[spir.Identifier, spir.IRType], + header_code: StringIO): + if len(statement.variables) != 2 or len(statement.range_expression) != 2: + return None + if len(statement.body) != 1 or not isinstance(statement.body[0], spir.AssignmentStatement): + return None + + # Loop-variable roles are detected from the destination index's affine decomposition rather + # than from statement.variables[0]/[1] position, so both declaration orders — for (k, l) and + # for (l, k) in [...] — are accepted: whichever variable appears alone with coefficient 1 in + # the destination index is the row variable (k); the other is the reduction variable (l). + var_names = [v.identifier.name for v in statement.variables] + var_by_name = {v.identifier.name: v for v in statement.variables} + range_by_name = dict(zip(var_names, statement.range_expression)) + + assign = statement.body[0] + dst = assign.destination + dst_idx = _single_index(dst) + if dst_idx is None: + return None + dst_aff = _affine_of(dst_idx, set(var_names)) + row_candidates = [name for name in var_names if dst_aff == ({name: 1}, 0)] + if len(row_candidates) != 1: + return None + k_var = row_candidates[0] + l_var = next(n for n in var_names if n != k_var) + + try: + rk, rl = range_by_name[k_var], range_by_name[l_var] + k0 = rk.start.eval() if rk.start is not None else 0 + kk = rk.stop.eval() + ks = rk.step.eval() if rk.step is not None else 1 + l0 = rl.start.eval() if rl.start is not None else 0 + ll = rl.stop.eval() + ls = rl.step.eval() if rl.step is not None else 1 + except Exception: + return None + if not all(isinstance(v, int) for v in (k0, kk, ks, l0, ll, ls)) or (k0, ks, ls) != (0, 1, 1): + return None + + src = assign.source.value + if isinstance(src, spir.MultiplyAccumulateOperator): + acc, mul_b, mul_c = src.a.value, src.b.value, src.c.value + elif isinstance(src, spir.BinaryOperator) and src.op == '+': + acc = src.left.value + mul = src.right.value + if not (isinstance(mul, spir.BinaryOperator) and mul.op == '*'): + acc, mul = src.right.value, src.left.value + if not (isinstance(mul, spir.BinaryOperator) and mul.op == '*'): + return None + mul_b, mul_c = mul.left.value, mul.right.value + else: + return None + + # 누산 항은 목적지와 동일한 Z[k] + if not (isinstance(acc, spir.ArraySlice) and acc.array.name == dst.array.name): + return None + acc_idx = _single_index(acc) + if acc_idx is None or _affine_of(acc_idx, {k_var}) != ({k_var: 1}, 0): + return None + + # 곱 항: 하나는 k·l 아핀 접근의 배열 A, 하나는 l 만의 스칼라 접근 X + def classify(node): + if not isinstance(node, spir.ArraySlice): + return None + idx = _single_index(node) + if idx is None: + return None + aff = _affine_of(idx, {k_var, l_var}) + if aff is None: + return None + coeffs, const = aff + if coeffs.get(k_var) and coeffs[k_var] > 0: + return ('mat', node, coeffs.get(k_var), coeffs.get(l_var, 0), const) + if k_var not in coeffs: + return ('vec', node, idx) + return None + + cb, cc = classify(mul_b), classify(mul_c) + if cb and cb[0] == 'mat' and cc and cc[0] == 'vec': + mat, vec = cb, cc + elif cc and cc[0] == 'mat' and cb and cb[0] == 'vec': + mat, vec = cc, cb + else: + return None + _, mat_node, ck, cl, c0 = mat + _, vec_node, vec_idx = vec + if _ids_in(vec_idx) - {l_var} != set(): + return None + + # Aliasing guard: the DSD op reorders reads relative to the sequential scalar loop, + # so reading the written array through A or X must not be vectorized. + if mat_node.array.name == dst.array.name or vec_node.array.name == dst.array.name: + return None + + # Type dispatch mirrors FMADSDOp._as_csl (dsd_ops.py): reuse its dtype resolution so a + # scalar/array/stream-typed identifier resolves to the same base ScalarType FMADSDOp itself + # would see. Both multiplicands must share a dtype, and only the f16/f32 combinations + # FMADSDOp models (@fmach, @fmachs, @fmacs) are supported -- integer types are out of scope + # because no integer MAC builtin exists in DSD_ASSIGNMENT_MAPPING. + mat_dtype = dsd_ops._get_base_dtype(dtypes, mat_node) + vec_dtype = dsd_ops._get_base_dtype(dtypes, vec_node) + dst_dtype = dsd_ops._get_base_dtype(dtypes, dst) + if mat_dtype != vec_dtype: + return None + mac_op = _mac_vectorization_op(mat_dtype, dst_dtype) + if mac_op is None: + return None + elem_type = dtype_as_csl(mat_dtype) + + uid = _VEC_COUNTER[0] + _VEC_COUNTER[0] += 1 + z_name = name_to_csl(dst.array) + a_name = name_to_csl(mat_node.array) + x_l = expr_to_csl(vec_idx) + dst_dsd = f'__vecmac_dst_{uid}' + src_dsd = f'__vecmac_src_{uid}' + base_expr = f'__index * {ck}' + (f' + {c0}' if c0 else '') + header_code.write( + f'const {dst_dsd} = @get_dsd(mem1d_dsd, .{{ .tensor_access = |__index|{{{kk}}} -> {z_name}[__index] }});\n' + f'const {src_dsd} = @get_dsd(mem1d_dsd, .{{ .tensor_access = |__index|{{{kk}}} -> {a_name}[{base_expr}] }});\n') + off = f'{name_to_csl(var_by_name[l_var].identifier)}' + (f' * {cl}' if cl != 1 else '') + lname = name_to_csl(var_by_name[l_var].identifier) + return ( + f'// [P4] vectorized MAC: {z_name}[k] += {a_name}[k*{ck}+{lname}*{cl}+{c0}] * {name_to_csl(vec_node.array)}[{x_l}]\n' + f'for (@range(i16, {l0}, {ll}, 1)) |{lname}| {{\n' + f' const __vecmac_a_{uid} = @increment_dsd_offset({src_dsd}, {off}, {elem_type});\n' + f' {mac_op}({dst_dsd}, {dst_dsd}, __vecmac_a_{uid}, {name_to_csl(vec_node.array)}[{x_l}]);\n' + f'}}\n') + + def emit_for(statement: spir.ForStatement, dsds: UniqueDSDDict, dtypes: dict[spir.Identifier, spir.IRType], header_code: StringIO) -> str: """ @@ -261,6 +484,11 @@ def emit_for(statement: spir.ForStatement, dsds: UniqueDSDDict, dtypes: dict[spi :param header_code: The header code to include. :return: The generated CSL for loop statement. """ + if not dsd_ops.DISABLE_DSD: + vectorized = _try_emit_vectorized_mac(statement, dtypes, header_code) + if vectorized is not None: + return vectorized + ranges = statement.range_expression vars_ = statement.variables diff --git a/tests/spatial_ir/test_mac_vectorization.py b/tests/spatial_ir/test_mac_vectorization.py new file mode 100644 index 00000000..6330ba06 --- /dev/null +++ b/tests/spatial_ir/test_mac_vectorization.py @@ -0,0 +1,195 @@ +"""Tests for the targeted vectorization of nested multiply-accumulate loops. + +The lowering should turn + + for (k, l) in [0:K, 0:K]: + z[k] = z[k] + A[k*K + l] * x[l] + +into a strided base DSD plus a per-column ``@increment_dsd_offset`` and one of +``@fmach``/``@fmachs``/``@fmacs`` (dtype-dependent, mirroring FMADSDOp), and must fall back to +scalar loops for every shape it cannot prove safe (aliasing, non-affine indices, unsupported +dtype combinations, or ``--disable-dsd``). There is no separate disable flag for this pass: it is +gated solely by ``dsd_ops.DISABLE_DSD``, same as every other DSD operation. +""" +import pytest +from spada.lowering.spatial_ir_to_csl import lower_spatial_ir_to_csl +from spada.syntax.csl import dsd_ops +from spada.syntax.spatial_ir import parser, passes + + +def _kernel(body: str, decls: str = None, mat_ty: str = 'f32', out_ty: str = None) -> str: + # mat_ty covers A_flat/x (the two @fmac* multiplicands, which must share a dtype); out_ty + # covers z (the accumulator/destination), defaulting to mat_ty. Stream types must match the + # local arrays they feed (receive/send are otherwise rejected as a dtype mismatch upstream of + # MAC vectorization entirely), so both need to vary together with the local declarations. + out_ty = out_ty if out_ty is not None else mat_ty + decls = decls if decls is not None else f''' + {mat_ty}[K*K] A_flat + {mat_ty}[K] x + {out_ty}[K] z + ''' + return f''' + kernel @t( + stream<{mat_ty}, K*K>[2, 2] readonly inp_A, + stream<{mat_ty}, K>[2, 2] readonly inp_x, + stream<{out_ty}, K>[2, 2] writeonly out + ) {{ + place i16 i, i16 j in [0:2, 0:2] {{ + {decls} + }} + phase {{ + compute i16 i, i16 j in [0:2, 0:2] {{ + await receive(A_flat, inp_A[i, j]) + await receive(x, inp_x[i, j]) + {body} + await send(z, out[i, j]) + }} + }} + }} + ''' + + +def _lower(code: str, **lower_kwargs): + kernel = parser.parse_string(code, 'test.sptl') + kernel = passes.concretize_parameters(kernel, K=4) + kernel = passes.constexpr_propagation(kernel) + return lower_spatial_ir_to_csl(kernel, **lower_kwargs) + + +def _all_code(csl_files) -> str: + return '\n'.join(f.code for f in csl_files) + + +MAC_BODY = ''' + for i16 k in [0:K] { + z[k] = 0.0 + } + for i16 k, i16 l in [0:K, 0:K] { + z[k] = z[k] + A_flat[k*K + l] * x[l] + } +''' + + +def test_mac_loop_is_vectorized(): + code = _all_code(_lower(_kernel(MAC_BODY))) + assert '@fmacs(' in code + assert '@increment_dsd_offset(' in code + + +def test_scalar_fallback_when_disabled(): + # MAC vectorization has no dedicated disable flag -- it is folded into --disable-dsd, the + # same switch that disables every other DSD operation (issue #69 ask 5). + try: + code = _all_code(_lower(_kernel(MAC_BODY), disable_dsd=True)) + assert '@increment_dsd_offset(' not in code + finally: + dsd_ops.DISABLE_DSD = False # module flag is sticky, reset for other tests + + +# --- Loop-variable-order matrix (issue #69 ask 2: (k, l) and (l, k) must both be detected) ----- + +def _order_decl(order: str) -> str: + """The `for in [0:K, 0:K]` variable declaration for a given loop-variable order.""" + return {'kl': 'i16 k, i16 l', 'lk': 'i16 l, i16 k'}[order] + + +def _mac_body(order: str) -> str: + return f''' + for i16 k in [0:K] {{ + z[k] = 0.0 + }} + for {_order_decl(order)} in [0:K, 0:K] {{ + z[k] = z[k] + A_flat[k*K + l] * x[l] + }} + ''' + + +ORDERS = ['kl', 'lk'] + +# (mat_ty, out_ty, expected op) for the dtype combinations FMADSDOp actually models +# (dsd_ops.DSD_ASSIGNMENT_MAPPING has no integer @fmac* builtin, so only f16/f32 are covered). +SUPPORTED_DTYPE_COMBOS = [ + ('f32', 'f32', '@fmacs('), # both multiplicands and accumulator f32 + ('f16', 'f16', '@fmach('), # both multiplicands and accumulator f16 + ('f32', 'f16', '@fmachs('), # f32 multiplicands, f16 accumulate (mixed) +] + +# (mat_ty, out_ty) combinations that must NOT vectorize: either FMADSDOp itself has no dispatch +# branch for the combination (f16 multiplicands / f32 accumulate), or no @fmac* builtin exists +# for the dtype at all (integer). +UNSUPPORTED_DTYPE_COMBOS = [ + ('f16', 'f32'), + ('i16', 'i16'), +] + + +@pytest.mark.parametrize('order', ORDERS) +@pytest.mark.parametrize('mat_ty, out_ty, expected_op', SUPPORTED_DTYPE_COMBOS) +def test_mac_loop_vectorized_for_order_and_dtype(order, mat_ty, out_ty, expected_op): + code = _all_code(_lower(_kernel(_mac_body(order), mat_ty=mat_ty, out_ty=out_ty))) + assert '@increment_dsd_offset(' in code + assert expected_op in code + + +@pytest.mark.parametrize('order', ORDERS) +@pytest.mark.parametrize('mat_ty, out_ty', UNSUPPORTED_DTYPE_COMBOS) +def test_no_vectorization_for_unsupported_dtype_combo(order, mat_ty, out_ty): + code = _all_code(_lower(_kernel(_mac_body(order), mat_ty=mat_ty, out_ty=out_ty))) + assert '@increment_dsd_offset(' not in code + + +@pytest.mark.parametrize('order', ORDERS) +def test_no_vectorization_when_accumulator_aliases_matrix(order): + # z appears as the matrix operand: DSD reordering would change semantics. + body = f''' + for i16 k in [0:K] {{ + z[k] = 1.0 + }} + for {_order_decl(order)} in [0:K, 0:K] {{ + z[k] = z[k] + z[k*1 + l*0] * x[l] + }} + ''' + code = _all_code(_lower(_kernel(body))) + assert '@increment_dsd_offset(' not in code + + +@pytest.mark.parametrize('order', ORDERS) +def test_no_vectorization_when_accumulator_aliases_vector(order): + body = f''' + for i16 k in [0:K] {{ + z[k] = 1.0 + }} + for {_order_decl(order)} in [0:K, 0:K] {{ + z[k] = z[k] + A_flat[k*K + l] * z[l] + }} + ''' + code = _all_code(_lower(_kernel(body))) + assert '@increment_dsd_offset(' not in code + + +@pytest.mark.parametrize('order', ORDERS) +def test_no_vectorization_for_nonaffine_index(order): + body = f''' + for i16 k in [0:K] {{ + z[k] = 0.0 + }} + for {_order_decl(order)} in [0:K, 0:K] {{ + z[k] = z[k] + A_flat[k*l] * x[l] + }} + ''' + code = _all_code(_lower(_kernel(body))) + assert '@increment_dsd_offset(' not in code + + +@pytest.mark.parametrize('order', ORDERS) +def test_no_vectorization_when_accumulator_differs(order): + body = f''' + for i16 k in [0:K] {{ + z[k] = 0.0 + }} + for {_order_decl(order)} in [0:K, 0:K] {{ + z[k] = x[k] + A_flat[k*K + l] * x[l] + }} + ''' + code = _all_code(_lower(_kernel(body))) + assert '@increment_dsd_offset(' not in code