pypto_pro.language.matmul_mx#

产品支持情况#

  • Ascend 950PR/Ascend 950DT:支持

  • Atlas A3 训练系列产品/Atlas A3 推理系列产品:不支持

  • Atlas A2 训练系列产品/Atlas A2 推理系列产品:不支持

功能说明#

完成带量化系数的矩阵乘法:

\[ C_{M \times N} = \left(\mathrm{scale\_a}_{M \times (K/32)} \otimes A_{M \times K}\right) \times \left(\mathrm{scale\_b}_{(K/32) \times N} \otimes B_{K \times N}\right) \]

其中,⊗表示广播乘法,A、B沿K维度每连续32个元素分别共享scale_a、scale_b中的一个量化系数。

函数原型#

pypto_pro.language.matmul_mx(
    dst_tile: Tile,
    lhs_tile: Tile,
    rhs_tile: Tile,
    scale_a: Tile,
    scale_b: Tile,
    *,
    phase: Optional[AccPhase] = None,
) -> None

参数说明#

参数

输入/输出

说明

dst_tile

输出

目的操作数,Tile类型,存储空间为L0C Buffer,shape为[M, N],数据类型为DT_FP32,layout必须为NZ。

lhs_tile

输入

源操作数(A,左矩阵),Tile类型,存储空间为L0A Buffer,数据类型支持DT_FP8E4M3FN、DT_FP8E5M2、DT_FP4E2M1和DT_FP4E1M2,layout必须为NZ,K必须为64的倍数。

rhs_tile

输入

源操作数(B,右矩阵),Tile类型,存储空间为L0B Buffer,数据类型支持DT_FP8E4M3FN、DT_FP8E5M2、DT_FP4E2M1和DT_FP4E1M2。与lhs_tile的数据类型可以不同,但必须同时选自DT_FP8E4M3FN、DT_FP8E5M2,或同时选自DT_FP4E2M1、DT_FP4E1M2。layout必须为ZN,K必须为64的倍数。

scale_a

输入

源操作数(左量化系数矩阵),Tile类型,存储空间为L0A_MX Buffer,数据类型为DT_FP8E8M0,shape为[M, K/32],layout默认且仅支持ZZ,fractal默认且仅支持32。每个量化系数对应A矩阵K方向连续32个元素。

scale_b

输入

源操作数(右量化系数矩阵),Tile类型,存储空间为L0B_MX Buffer,数据类型为DT_FP8E8M0,shape为[K/32, N],layout默认且仅支持NN,fractal默认且仅支持32。每个量化系数对应B矩阵K方向连续32个元素。

phase

输入

K维分块累加阶段,pypto_pro.language.AccPhase类型,可选,用于控制矩阵计算与L0C Buffer数据搬出之间的UnitFlag同步。与pypto_pro.language.STPhase的配合方式见AccPhase与STPhase配合使用说明

约束说明#

  • 量化系数Tile与A/B矩阵Tile必须满足以下硬件地址关系:

    addr(scale_a) = addr(lhs_tile) >> 4
    addr(scale_b) = addr(rhs_tile) >> 4
    

    硬件根据A/B矩阵首地址定位scale_a/scale_b,地址不满足上述关系时计算结果可能错误。

返回值说明#

无。

调用示例#

MXFP8矩阵乘#

下面计算C[128, 128] = A[128, 128] @ B[128, 128]。A和B的K方向均包含4个量化系数分组,因此scale_a/scale_b Tile的逻辑shape分别为[128, 4]和[4, 128],GM中的Tensor物理shape分别为[128, 2, 2]和[2, 128, 2]。

import os
import pypto_pro.language as pl
import torch

TILE = 128
SCALE_K = TILE // 32


@pl.jit(auto_mutex=True)
def mxfp8_matmul_kernel(
    a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP8E4M3FN],
    b: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP8E5M2],
    scale_a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC, 2], pl.DT_FP8E8M0],
    scale_b: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC, 2], pl.DT_FP8E8M0],
    out: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32],
):
    a_l1 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP8E4M3FN,
                         target_memory=pl.MemorySpace.Mat, layout=pl.NZ),
        addrs=0x00000, mutex_ids=[0])
    b_l1 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP8E5M2,
                         target_memory=pl.MemorySpace.Mat, layout=pl.NZ),
        addrs=0x10000, mutex_ids=[1])
    scale_a_l1 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, SCALE_K], dtype=pl.DT_FP8E8M0,
                         target_memory=pl.MemorySpace.Mat, layout=pl.ZZ),
        addrs=0x20000, mutex_ids=[2])
    scale_b_l1 = pl.make_tile_group(
        type=pl.TileType(shape=[SCALE_K, TILE], dtype=pl.DT_FP8E8M0,
                         target_memory=pl.MemorySpace.Mat, layout=pl.NN),
        addrs=0x21000, mutex_ids=[3])
    a_l0 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP8E4M3FN,
                         target_memory=pl.MemorySpace.Left, layout=pl.NZ),
        addrs=0x0000, mutex_ids=[4])
    b_l0 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP8E5M2,
                         target_memory=pl.MemorySpace.Right, layout=pl.ZN),
        addrs=0x0000, mutex_ids=[5])
    scale_a_l0 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, SCALE_K], dtype=pl.DT_FP8E8M0,
                         target_memory=pl.MemorySpace.ScaleLeft, layout=pl.ZZ),
        addrs=0x0000, mutex_ids=[6])
    scale_b_l0 = pl.make_tile_group(
        type=pl.TileType(shape=[SCALE_K, TILE], dtype=pl.DT_FP8E8M0,
                         target_memory=pl.MemorySpace.ScaleRight, layout=pl.NN),
        addrs=0x0000, mutex_ids=[7])
    acc = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP32,
                         target_memory=pl.MemorySpace.Acc, layout=pl.NZ, fractal=1024),
        addrs=0x0000, mutex_ids=[8])

    with pl.section_cube():
        al1, bl1 = a_l1.current(), b_l1.current()
        sal1, sbl1 = scale_a_l1.current(), scale_b_l1.current()
        al0, bl0 = a_l0.current(), b_l0.current()
        sal0, sbl0 = scale_a_l0.current(), scale_b_l0.current()
        ac = acc.current()
        pl.load(al1, a, [0, 0])
        pl.load(bl1, b, [0, 0])
        pl.load(sal1, scale_a, [0, 0, 0], order=[0, 1])
        pl.load(sbl1, scale_b, [0, 0, 0], order=[0, 1])
        pl.move(al0, al1)
        pl.move(bl0, bl1)
        pl.move(sal0, sal1)
        pl.move(sbl0, sbl1)
        pl.matmul_mx(ac, al0, bl0, sal0, sbl0)
        pl.store(out, ac, [0, 0])


if __name__ == "__main__":
    device = f"npu:{int(os.environ.get('TILE_FWK_DEVICE_ID', 0))}"
    torch.npu.set_device(device)

    # E4M3编码0x38、E5M2编码0x3C均表示1.0;E8M0编码127表示缩放值1.0。
    a = torch.full([TILE, TILE], 0x38, dtype=torch.uint8).view(torch.float8_e4m3fn).to(device)
    b = torch.full([TILE, TILE], 0x3C, dtype=torch.uint8).view(torch.float8_e5m2).to(device)
    scale_a = torch.full([TILE, SCALE_K // 2, 2], 127, device=device, dtype=torch.uint8)
    scale_b = torch.full([SCALE_K // 2, TILE, 2], 127, device=device, dtype=torch.uint8)
    out = torch.zeros([TILE, TILE], device=device, dtype=torch.float32)

    mxfp8_matmul_kernel(a, b, scale_a, scale_b, out)
    torch.npu.synchronize()

    ref = torch.full([TILE, TILE], float(TILE), device=device, dtype=torch.float32)
    torch.testing.assert_close(out, ref, rtol=1e-3, atol=1e-3)
    print(f"max diff = {(out - ref).abs().max().item()}")

MXFP4矩阵乘#

下面计算C[128, 128] = A[128, 128] @ B[128, 128]。A和B的K方向均包含4个量化系数分组,因此scale_a/scale_b Tile的逻辑shape分别为[128, 4]和[4, 128],GM中的Tensor物理shape分别为[128, 2, 2]和[2, 128, 2]。FP4矩阵在运行时以uint8存储,每字节打包两个FP4元素。

import os
import pypto_pro.language as pl
import torch

TILE = 128
SCALE_K = TILE // 32


@pl.jit(auto_mutex=True)
def mxfp4_matmul_kernel(
    a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP4E2M1],
    b: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP4E1M2],
    scale_a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC, 2], pl.DT_FP8E8M0],
    scale_b: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC, 2], pl.DT_FP8E8M0],
    out: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32],
):
    a_l1 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP4E2M1,
                         target_memory=pl.MemorySpace.Mat, layout=pl.NZ),
        addrs=0x00000, mutex_ids=[0])
    b_l1 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP4E1M2,
                         target_memory=pl.MemorySpace.Mat, layout=pl.NZ),
        addrs=0x10000, mutex_ids=[1])
    scale_a_l1 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, SCALE_K], dtype=pl.DT_FP8E8M0,
                         target_memory=pl.MemorySpace.Mat, layout=pl.ZZ),
        addrs=0x20000, mutex_ids=[2])
    scale_b_l1 = pl.make_tile_group(
        type=pl.TileType(shape=[SCALE_K, TILE], dtype=pl.DT_FP8E8M0,
                         target_memory=pl.MemorySpace.Mat, layout=pl.NN),
        addrs=0x21000, mutex_ids=[3])
    a_l0 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP4E2M1,
                         target_memory=pl.MemorySpace.Left, layout=pl.NZ),
        addrs=0x0000, mutex_ids=[4])
    b_l0 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP4E1M2,
                         target_memory=pl.MemorySpace.Right, layout=pl.ZN),
        addrs=0x0000, mutex_ids=[5])
    scale_a_l0 = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, SCALE_K], dtype=pl.DT_FP8E8M0,
                         target_memory=pl.MemorySpace.ScaleLeft, layout=pl.ZZ),
        addrs=0x0000, mutex_ids=[6])
    scale_b_l0 = pl.make_tile_group(
        type=pl.TileType(shape=[SCALE_K, TILE], dtype=pl.DT_FP8E8M0,
                         target_memory=pl.MemorySpace.ScaleRight, layout=pl.NN),
        addrs=0x0000, mutex_ids=[7])
    acc = pl.make_tile_group(
        type=pl.TileType(shape=[TILE, TILE], dtype=pl.DT_FP32,
                         target_memory=pl.MemorySpace.Acc, layout=pl.NZ, fractal=1024),
        addrs=0x0000, mutex_ids=[8])

    with pl.section_cube():
        al1, bl1 = a_l1.current(), b_l1.current()
        sal1, sbl1 = scale_a_l1.current(), scale_b_l1.current()
        al0, bl0 = a_l0.current(), b_l0.current()
        sal0, sbl0 = scale_a_l0.current(), scale_b_l0.current()
        ac = acc.current()
        pl.load(al1, a, [0, 0])
        pl.load(bl1, b, [0, 0])
        pl.load(sal1, scale_a, [0, 0, 0], order=[0, 1])
        pl.load(sbl1, scale_b, [0, 0, 0], order=[0, 1])
        pl.move(al0, al1)
        pl.move(bl0, bl1)
        pl.move(sal0, sal1)
        pl.move(sbl0, sbl1)
        pl.matmul_mx(ac, al0, bl0, sal0, sbl0)
        pl.store(out, ac, [0, 0])


def pack_fp4(codes: torch.Tensor) -> torch.Tensor:
    if codes.shape[-1] % 2 != 0:
        raise ValueError("FP4 packing requires an even last dimension")
    low = codes[..., 0::2] & 0x0F
    high = (codes[..., 1::2] & 0x0F) << 4
    return (low | high).contiguous()


if __name__ == "__main__":
    device = f"npu:{int(os.environ.get('TILE_FWK_DEVICE_ID', 0))}"
    torch.npu.set_device(device)

    # E2M1编码0x2、E1M2编码0x4均表示1.0;E8M0编码127表示缩放值1.0。
    a_codes = torch.full([TILE, TILE], 0x2, dtype=torch.uint8)
    b_codes = torch.full([TILE, TILE], 0x4, dtype=torch.uint8)
    a = pack_fp4(a_codes).to(device)
    b = pack_fp4(b_codes).to(device)
    scale_a = torch.full([TILE, SCALE_K // 2, 2], 127, device=device, dtype=torch.uint8)
    scale_b = torch.full([SCALE_K // 2, TILE, 2], 127, device=device, dtype=torch.uint8)
    out = torch.zeros([TILE, TILE], device=device, dtype=torch.float32)

    mxfp4_matmul_kernel(a, b, scale_a, scale_b, out)
    torch.npu.synchronize()

    ref = torch.full([TILE, TILE], float(TILE), device=device, dtype=torch.float32)
    torch.testing.assert_close(out, ref, rtol=1e-3, atol=1e-3)
    print(f"max diff = {(out - ref).abs().max().item()}")