|
22 | 22 | wave_compile, |
23 | 23 | ) |
24 | 24 | from wave_lang.kernel.wave.type_inference import infer_types |
25 | | -from wave_lang.kernel.wave.constraints import Constraint, MMAType |
| 25 | +from wave_lang.kernel.wave.constraints import Constraint, MMAType, ScaledMMAType |
26 | 26 | from wave_lang.kernel.wave.mlir_converter.mlir_converter import ( |
27 | 27 | format_diagnostics, |
28 | 28 | PersistentEmitter, |
@@ -1866,3 +1866,139 @@ def mlir_converter_attention_pre_infer_types(): |
1866 | 1866 |
|
1867 | 1867 | # Iterate result types match init_arg types. |
1868 | 1868 | # CHECK: -> (!wave.tensor<[@B, @M] of f32, <register>>, !wave.tensor<[@B, @M] of f32, <register>>, !wave.tensor<[@B, @N, @M] of f32, <register>>) |
| 1869 | + |
| 1870 | + |
| 1871 | +@run_test |
| 1872 | +def test_mxfp4_scaled_mma_256x256x256(): |
| 1873 | + # Input sizes |
| 1874 | + M = tkl.sym.M |
| 1875 | + N = tkl.sym.N |
| 1876 | + K = tkl.sym.K |
| 1877 | + # Workgroup tile sizes |
| 1878 | + BLOCK_M = tkl.sym.BLOCK_M |
| 1879 | + BLOCK_N = tkl.sym.BLOCK_N |
| 1880 | + BLOCK_K = tkl.sym.BLOCK_K |
| 1881 | + # Address space (for GPU, shared(1) or global(0)) |
| 1882 | + ADDRESS_SPACE = tkl.sym.ADDRESS_SPACE |
| 1883 | + |
| 1884 | + mfma_variant = ScaledMMAType.F32_16x16x128_F8F6F4 |
| 1885 | + |
| 1886 | + # Expose user-constraints |
| 1887 | + constraints: list[tkw.Constraint] = [tkw.WorkgroupConstraint(M, BLOCK_M, 0)] |
| 1888 | + constraints += [tkw.WorkgroupConstraint(N, BLOCK_N, 1)] |
| 1889 | + constraints += [tkw.TilingConstraint(K, BLOCK_K)] |
| 1890 | + constraints += [tkw.WaveConstraint(M, BLOCK_M / 2)] |
| 1891 | + constraints += [tkw.WaveConstraint(N, BLOCK_N / 2)] |
| 1892 | + |
| 1893 | + constraints += [tkw.HardwareConstraint(threads_per_wave=64, mma_type=mfma_variant)] |
| 1894 | + |
| 1895 | + @tkw.wave(constraints) |
| 1896 | + def scaled_gemm( |
| 1897 | + a: tkl.Memory[M, K / 2, ADDRESS_SPACE, tkl.i8], |
| 1898 | + a_scale: tkl.Memory[M, K / 32, ADDRESS_SPACE, tkl.i8], |
| 1899 | + b: tkl.Memory[N, K / 2, ADDRESS_SPACE, tkl.i8], |
| 1900 | + b_scale: tkl.Memory[N, K / 32, ADDRESS_SPACE, tkl.i8], |
| 1901 | + c: tkl.Memory[M, N, GLOBAL_ADDRESS_SPACE, tkl.f32], |
| 1902 | + ): |
| 1903 | + c_reg = tkl.Register[M, N, tkl.f32](0.0) |
| 1904 | + |
| 1905 | + @tkw.iterate(K, init_args=[c_reg]) |
| 1906 | + def repeat(acc: tkl.Register[M, N, tkl.f32]) -> tkl.Register[M, N, tkl.f32]: |
| 1907 | + a_reg = tkw.read(a) |
| 1908 | + a_reg = tkw.bitcast(a_reg, tkl.f4e2m1fn) |
| 1909 | + a_scale_reg = tkw.read(a_scale) |
| 1910 | + a_scale_reg = tkw.bitcast(a_scale_reg, tkl.f8e8m0fnu) |
| 1911 | + b_reg = tkw.read(b) |
| 1912 | + b_reg = tkw.bitcast(b_reg, tkl.f4e2m1fn) |
| 1913 | + b_scale_reg = tkw.read(b_scale) |
| 1914 | + b_scale_reg = tkw.bitcast(b_scale_reg, tkl.f8e8m0fnu) |
| 1915 | + acc = tkw.scaled_mma(a_reg, a_scale_reg, b_reg, b_scale_reg, acc) |
| 1916 | + return acc |
| 1917 | + |
| 1918 | + tkw.write(repeat, c) |
| 1919 | + |
| 1920 | + subs = { |
| 1921 | + ADDRESS_SPACE: SHARED_ADDRESS_SPACE, |
| 1922 | + BLOCK_M: 32, |
| 1923 | + BLOCK_N: 32, |
| 1924 | + BLOCK_K: 256, |
| 1925 | + M: 16384, |
| 1926 | + N: 16384, |
| 1927 | + K: 16384, |
| 1928 | + } |
| 1929 | + options = WaveCompileOptions( |
| 1930 | + subs=subs, |
| 1931 | + compile_to_mlir=True, |
| 1932 | + location_capture_config=LocationCaptureConfig(level=LocationCaptureLevel.NONE), |
| 1933 | + enforce_locations=False, |
| 1934 | + ) |
| 1935 | + options = set_default_run_config(options) |
| 1936 | + |
| 1937 | + compiled_kernel = wave_compile(options, scaled_gemm) |
| 1938 | + trace = compiled_kernel.get_compiled_graph() |
| 1939 | + |
| 1940 | + mlir_output, diagnostics, _ = emitter.emit_wave_dialect( |
| 1941 | + trace, scaled_gemm.constraints, options |
| 1942 | + ) |
| 1943 | + |
| 1944 | + if diagnostics: |
| 1945 | + print(format_diagnostics(diagnostics, use_color=False), file=sys.stderr) |
| 1946 | + assert ( |
| 1947 | + len(diagnostics) == 0 |
| 1948 | + ), "dialect emission should create valid IR, therefore diagnostics should be empty" |
| 1949 | + |
| 1950 | + print(mlir_output) |
| 1951 | + |
| 1952 | + # CHECK-LABEL: test_mxfp4_scaled_mma_256x256x256 |
| 1953 | + # CHECK: module |
| 1954 | + # CHECK-NEXT: func.func @kernel( |
| 1955 | + # CHECK-SAME: %[[A:.*]]: !wave.tensor<[@M, @K2] of i8, <global>> |
| 1956 | + # CHECK-SAME: %[[A_SCALE:.*]]: !wave.tensor<[@M, @K32] of i8, <global>> |
| 1957 | + # CHECK-SAME: %[[B:.*]]: !wave.tensor<[@N, @K2] of i8, <global>> |
| 1958 | + # CHECK-SAME: %[[B_SCALE:.*]]: !wave.tensor<[@N, @K32] of i8, <global>> |
| 1959 | + # CHECK-SAME: %[[C:.*]]: !wave.tensor<[@M, @N] of f32, <global>> |
| 1960 | + # CHECK-SAME: wave.constraints = |
| 1961 | + # CHECK-SAME: #wave.workgroup_constraint<dim = <"M">, tile_size = <[#wave.symbol<"BLOCK_M">] -> (BLOCK_M)>, workgroup_dim = <x>> |
| 1962 | + # CHECK-SAME: #wave.workgroup_constraint<dim = <"N">, tile_size = <[#wave.symbol<"BLOCK_N">] -> (BLOCK_N)>, workgroup_dim = <y>> |
| 1963 | + # CHECK-SAME: #wave.tiling_constraint<dim = <"K">, tile_size = <[#wave.symbol<"BLOCK_K">] -> (BLOCK_K)>> |
| 1964 | + # CHECK-SAME: #wave.wave_constraint<dim = <"M">, tile_size = <[#wave.symbol<"BLOCK_M">] -> (BLOCK_M floordiv 2)>> |
| 1965 | + # CHECK-SAME: #wave.wave_constraint<dim = <"N">, tile_size = <[#wave.symbol<"BLOCK_N">] -> (BLOCK_N floordiv 2)>> |
| 1966 | + # CHECK-SAME: #wave.hardware_constraint<threads_per_wave = 64, waves_per_block = [2, 2, 1], mma_type = <f32_16x16x128_f8f6f4>> |
| 1967 | + # CHECK-SAME: #wave.hyperparameters<{BLOCK_K = 256 : i64, BLOCK_M = 32 : i64, BLOCK_N = 32 : i64, K = 16384 : i64, K2 = #wave.expr_list<[#wave.symbol<"K">] -> (K floordiv 2)>, K32 = #wave.expr_list<[#wave.symbol<"K">] -> (K floordiv 32)>, M = 16384 : i64, N = 16384 : i64}> |
| 1968 | + # |
| 1969 | + # CHECK: %[[CST:.*]] = arith.constant 0.000000e+00 : f32 |
| 1970 | + # CHECK-NEXT: %[[REG:.*]] = wave.register %[[CST]] |
| 1971 | + # |
| 1972 | + # CHECK: %[[ITERATE:.*]] = wave.iterate @K iter_args(%[[REG]]) { |
| 1973 | + # CHECK-NEXT: ^{{.*}}(%[[ACC:.*]]: !wave.tensor<[@M, @N] of f32, <register>>): |
| 1974 | + # |
| 1975 | + # Global reads promoted through shared memory. |
| 1976 | + # |
| 1977 | + # CHECK: wave.read %[[A]] |
| 1978 | + # CHECK: wave.write {{.*}} !wave.tensor<[@M, @K2] of i8, <shared>> |
| 1979 | + # CHECK: wave.read %[[A_SCALE]] |
| 1980 | + # CHECK: wave.write {{.*}} !wave.tensor<[@M, @K32] of i8, <shared>> |
| 1981 | + # CHECK: wave.read %[[B]] |
| 1982 | + # CHECK: wave.write {{.*}} !wave.tensor<[@N, @K2] of i8, <shared>> |
| 1983 | + # CHECK: wave.read %[[B_SCALE]] |
| 1984 | + # CHECK: wave.write {{.*}} !wave.tensor<[@N, @K32] of i8, <shared>> |
| 1985 | + # |
| 1986 | + # Bitcasts from i8 to scaled MMA operand types. |
| 1987 | + # |
| 1988 | + # CHECK: wave.bitcast {{.*}} : !wave.tensor<[@M, @K2] of i8, <register>> to !wave.tensor<[@M, @K] of f4E2M1FN, <register>> |
| 1989 | + # CHECK: wave.bitcast {{.*}} : !wave.tensor<[@M, @K32] of i8, <register>> to !wave.tensor<[@M, @K32] of f8E8M0FNU, <register>> |
| 1990 | + # CHECK: wave.bitcast {{.*}} : !wave.tensor<[@N, @K2] of i8, <register>> to !wave.tensor<[@N, @K] of f4E2M1FN, <register>> |
| 1991 | + # CHECK: wave.bitcast {{.*}} : !wave.tensor<[@N, @K32] of i8, <register>> to !wave.tensor<[@N, @K32] of f8E8M0FNU, <register>> |
| 1992 | + # |
| 1993 | + # Two scaled MMAs (BLOCK_K=256 with 128 per intrinsic). |
| 1994 | + # |
| 1995 | + # CHECK: %[[SMMA0:.*]] = wave.scaled_mma {{.*}}, {{.*}}, {{.*}}, {{.*}}, %[[ACC]] |
| 1996 | + # CHECK-SAME: #wave.mma_kind<f32_16x16x128_f8f6f4> |
| 1997 | + # CHECK: %[[SMMA1:.*]] = wave.scaled_mma {{.*}}, {{.*}}, {{.*}}, {{.*}}, %[[SMMA0]] |
| 1998 | + # CHECK-SAME: #wave.mma_kind<f32_16x16x128_f8f6f4> |
| 1999 | + # CHECK: wave.yield %[[SMMA1]] : !wave.tensor<[@M, @N] of f32, <register>> |
| 2000 | + # CHECK-NEXT: } |
| 2001 | + # |
| 2002 | + # Results written back to the output tensor. |
| 2003 | + # |
| 2004 | + # CHECK: wave.write {{.*}}, %[[C]] |
0 commit comments