Repository navigation
Vectorize nested multiply-accumulate loops into @fmac* DSD operations #72
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||
|---|---|---|---|---|---|---|
|
|
@@ -250,17 +250,265 @@ def emit_assignment(statement: spir.AssignmentStatement, dsds: UniqueDSDDict, dt | |||||
| return dsd_ops.DSD_ASSIGNMENT_MAPPING[dsd_op]() | ||||||
|
|
||||||
|
|
||||||
| # Targeted vectorization of the nested multiply-accumulate loop | ||||||
| # for (k, l) or (l, k) in either order: Z[k] = Z[k] + A[k*Ck + l*Cl + C0] * X[f(l)] | ||||||
| # into the handwritten CSL idiom (strided base DSD + per-l @increment_dsd_offset + @fmac*). | ||||||
| # The accumulator index decides the roles: the loop variable indexing Z is k, the other | ||||||
| # is the reduction variable l. Dtypes dispatch like FMADSDOp._as_csl (@fmach / @fmachs / | ||||||
| # @fmacs). Conservative: fires only when Z is distinct from A and X; anything else falls | ||||||
| # through to scalar loops. Gated behind --disable-dsd like all other DSD detection. | ||||||
| # @increment_dsd_offset is available from SDK 1.x (used by Cerebras' own csl-examples | ||||||
| # v1.4.0 cholesky benchmark). | ||||||
|
|
||||||
|
|
||||||
| def _affine_of(expr: spir.Expression, varnames: set[str]) -> Optional[tuple[dict[str, int], int]]: | ||||||
| """ | ||||||
| Decomposes an index expression into an affine form over the given loop variables. | ||||||
|
|
||||||
| :param expr: The index expression to decompose. | ||||||
| :param varnames: The loop-variable names that may carry coefficients. | ||||||
| :return: A ``(coefficients, constant)`` pair such that the expression equals | ||||||
| ``sum(coefficients[var] * var) + constant`` with integer coefficients, | ||||||
| or None if the expression is not affine in the given variables. | ||||||
| """ | ||||||
| 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 (ValueError, TypeError): | ||||||
| return None | ||||||
| return ({}, int(v)) if isinstance(v, int) else None | ||||||
| if isinstance(e, spir.UnaryOperator) and e.op == '-': | ||||||
| sub = _affine_of(e.value, varnames) | ||||||
| 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]: # left operand is a pure constant | ||||||
| s = left[1] | ||||||
| return {k: c * s for k, c in right[0].items()}, right[1] * s | ||||||
| if not right[0]: # right operand is a pure constant | ||||||
| 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: spir.Expression) -> set[str]: | ||||||
| """ | ||||||
| Collects the names of all identifiers appearing in an expression. | ||||||
|
|
||||||
| :param expr: The expression to walk. | ||||||
| :return: The set of identifier names. | ||||||
| """ | ||||||
| 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: spir.SpatialNode) -> Optional[spir.Expression]: | ||||||
| """ | ||||||
| Extracts the index of a one-dimensional single-element array access. | ||||||
|
|
||||||
| :param node: The AST node to inspect. | ||||||
| :return: The index expression if ``node`` is an ArraySlice with exactly one | ||||||
| non-range index, otherwise 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 _try_emit_vectorized_mac(statement: spir.ForStatement, dtypes: dict[spir.Identifier, spir.IRType], | ||||||
| header_code: StringIO) -> Optional[str]: | ||||||
| """ | ||||||
| Attempts to vectorize a two-variable nested multiply-accumulate loop into DSD operations. | ||||||
|
|
||||||
| Recognizes ``Z[k] = Z[k] + A[affine(k, l)] * X[f(l)]`` (in either loop-header order, | ||||||
| and either operand order) and emits a strided base DSD over A plus a per-``l`` | ||||||
| ``@increment_dsd_offset`` and the dtype-matched ``@fmac*`` builtin. The loop variable | ||||||
| that indexes the destination with coefficient 1 takes the ``k`` role; the other is the | ||||||
| reduction variable ``l``. Falls back for anything it cannot prove safe: aliasing of Z | ||||||
| with A or X, non-affine indices, unsupported dtype combinations, non-unit steps, or a | ||||||
| nonzero ``k`` start. | ||||||
|
|
||||||
| :param statement: The for-loop statement to inspect. | ||||||
| :param dtypes: The data types dictionary. | ||||||
| :param header_code: The header code to write the DSD declarations into. | ||||||
| :return: The generated CSL replacing the loop nest, or None to fall back to scalar code. | ||||||
| """ | ||||||
| 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 | ||||||
| names = [v.identifier.name for v in statement.variables] | ||||||
|
|
||||||
| # Role assignment: the loop variable that alone indexes the destination Z with | ||||||
| # coefficient 1 is k; the other is the reduction variable l. The loop-header | ||||||
| # order (k, l) vs (l, k) is irrelevant. | ||||||
| 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(names)) | ||||||
| if dst_aff is None or dst_aff[1] != 0 or len(dst_aff[0]) != 1: | ||||||
| return None | ||||||
| (k_var, k_coeff), = dst_aff[0].items() | ||||||
| if k_coeff != 1: | ||||||
| return None | ||||||
| k_pos = names.index(k_var) | ||||||
| l_pos = 1 - k_pos | ||||||
| l_var = names[l_pos] | ||||||
|
|
||||||
| def _range_consts(r): | ||||||
| start = r.start.eval() if r.start is not None else 0 | ||||||
| stop = r.stop.eval() | ||||||
| step = r.step.eval() if r.step is not None else 1 | ||||||
| if not all(isinstance(v, int) for v in (start, stop, step)): | ||||||
| raise TypeError | ||||||
| return start, stop, step | ||||||
| try: | ||||||
| k0, kk, ks = _range_consts(statement.range_expression[k_pos]) | ||||||
| l0, ll, ls = _range_consts(statement.range_expression[l_pos]) | ||||||
| except (ValueError, TypeError, AttributeError): | ||||||
| return None | ||||||
| if (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 | ||||||
|
|
||||||
| # The accumulator term must be the same Z[k] as the destination. | ||||||
| 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 | ||||||
|
|
||||||
| # Multiply operands: one is the matrix A accessed by an affine index in k and l, | ||||||
| # the other a scalar access X indexed by l only. | ||||||
| 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 | ||||||
|
|
||||||
| # Dtype dispatch — follows the mapping of FMADSDOp._as_csl exactly | ||||||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. the classification here repeats the already-existing code. |
||||||
| # (accumulator = Z, DSD multiplicand = A, scalar = X). | ||||||
| def base_dtype(name_node): | ||||||
| t = dtypes.get(name_node.array) | ||||||
| return getattr(t, 'base_type', None) or getattr(t, 'element_type', None) | ||||||
| z_t, a_t, x_t = base_dtype(dst), base_dtype(mat_node), base_dtype(vec_node) | ||||||
| if z_t == a_t == spir.ScalarType.f16 and x_t == spir.ScalarType.f16: | ||||||
| fmac_builtin, elem_type = '@fmach', 'f16' | ||||||
| elif z_t == a_t == spir.ScalarType.f32 and x_t == spir.ScalarType.f16: | ||||||
| fmac_builtin, elem_type = '@fmachs', 'f32' # 16-bit multiplication, 32-bit addition | ||||||
| elif z_t == a_t == spir.ScalarType.f32 and x_t == spir.ScalarType.f32: | ||||||
| fmac_builtin, elem_type = '@fmacs', 'f32' | ||||||
| else: | ||||||
| return None | ||||||
|
|
||||||
| # Deterministic per-rectangle numbering: each vectorized loop writes exactly one | ||||||
| # __vecmac_dst_ declaration into this rectangle's header. | ||||||
| uid = header_code.getvalue().count('__vecmac_dst_') | ||||||
| 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, ' | ||||||
| f'.{{ .tensor_access = |__index|{{{kk}}} -> {z_name}[__index] }});\n' | ||||||
| f'const {src_dsd} = @get_dsd(mem1d_dsd, ' | ||||||
| f'.{{ .tensor_access = |__index|{{{kk}}} -> {a_name}[{base_expr}] }});\n') | ||||||
|
Comment on lines
+475
to
+482
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This is strange, there is already existing code that sets names for the DSD ops in the lowering. Why make new DSD names? |
||||||
| lname = name_to_csl(statement.variables[l_pos].identifier) | ||||||
| off = lname + (f' * {cl}' if cl != 1 else '') | ||||||
| return ( | ||||||
| f'// vectorized MAC: {z_name}[k] += {a_name}[k*{ck}+{lname}*{cl}+{c0}] * {name_to_csl(vec_node.array)}[{x_l}]\n' | ||||||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Very specific comment |
||||||
| f'for (@range(i16, {l0}, {ll}, 1)) |{lname}| {{\n' | ||||||
| f' const __vecmac_a_{uid} = @increment_dsd_offset({src_dsd}, {off}, {elem_type});\n' | ||||||
| f' {fmac_builtin}({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: | ||||||
| """ | ||||||
| Generates a CSL for loop statement from a Spatial IR for loop statement. | ||||||
|
|
||||||
| Nested multiply-accumulate loops are first offered to the DSD vectorizer (see | ||||||
| ``_try_emit_vectorized_mac``); everything else lowers to scalar loops. | ||||||
|
|
||||||
|
Comment on lines
+498
to
+500
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
Suggested change
|
||||||
| :param statement: The Spatial IR for loop statement to convert. | ||||||
| :param dsds: The unique DSD dictionary. | ||||||
| :param dtypes: The data types dictionary. | ||||||
| :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 | ||||||
|
Comment on lines
+507
to
+510
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Why is this overriding the entire statement generation instead of extending it? |
||||||
|
|
||||||
| ranges = statement.range_expression | ||||||
| vars_ = statement.variables | ||||||
|
|
||||||
|
|
||||||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
The code and comments is a bit too specialized on the k/l loops and does not belong as-is in a general language framework. Based on the helper functions, this can and should be generalized as per my comment.