SEP = "=" * 78
DASH = "-" * 78
fa_base = 0
fd_base = _AIC_CORE_NUM * _FA_META_SIZE
print(f"\n{DASH}")
print(f"FA Metadata — AIC cores "
f"(36 cores × 16 slots, 9 fields used)")
print(DASH)
cw = 15
hdr = f" {'Core':<7}" + "".join(f"{n:>{cw}}" for n in _FA_FIELDS)
print(hdr)
print(f" {'':<7}" + "".join(f"{'['+str(i)+']':>{cw}}" for i in range(len(_FA_FIELDS))))
print(f" {'─'*7}" + "─" * (cw * len(_FA_FIELDS)))
for core in range(_AIC_CORE_NUM):
off = core * _FA_META_SIZE
vals = [int(arr[off + i]) for i in range(len(_FA_FIELDS))]
note = " (inactive)" if all(v == 0 for v in vals) else ""
print(f" AIC{core:02d} " + "".join(f"{v:>{cw}}" for v in vals) + note)
print(f"\n{DASH}")
print(f"FD Metadata — active AIV cores only "
f"(M_NUM > 0 shown, 72 cores × 16 slots, 6 fields used)")
print(DASH)
cw = 12
hdr = f" {'Core':<7}" + "".join(f"{n:>{cw}}" for n in _FD_FIELDS)
print(hdr)
print(f" {'':<7}" + "".join(f"{'['+str(i)+']':>{cw}}" for i in range(len(_FD_FIELDS))))
print(f" {'─'*7}" + "─" * (cw * len(_FD_FIELDS)))
active = 0
for core in range(_AIV_CORE_NUM):
off = fd_base + core * _FD_META_SIZE
vals = [int(arr[off + i]) for i in range(len(_FD_FIELDS))]
if vals[5] == 0:
continue
print(f" AIV{core:02d} " + "".join(f"{v:>{cw}}" for v in vals))
active += 1
if active == 0:
print(" (no active FD cores)")
print(f"{SEP}\n")
import os
import sys
import torch
import torch_npu
import numpy as np
import torchair
from torchair.configs.compiler_config import CompilerConfig
from cann_ops_transformer.ops import sparse_flash_mla_metadata
_AIC_CORE_NUM = 36
_AIV_CORE_NUM = 72
_FA_META_SIZE = 9
_FD_META_SIZE = 8
_FA_FIELDS = ["CORE_ENABLE",
"BN2_START", "M_START", "S2_START",
"BN2_END", "M_END", "S2_END",
"FIRST_FD_WS_IDX", "S2_MAX_NUM"]
_FD_FIELDS = ["CORE_ENABLE",
"BN2_IDX", "M_IDX", "WS_IDX",
"WS_NUM", "M_START", "M_NUM"]
def print_metadata(metadata):
arr = metadata.cpu().contiguous().to(torch.int32).numpy().astype(np.uint32)
TEST_PARAMS = {
"HCA_decode_pa2": {
"testcase_name": ["csa_decode_pa"],
"layout_q": ["TND"],
"layout_kv": ["PA_BBND"],
"q_type": [torch.bfloat16],
"ori_kv_type": [torch.bfloat16],
"cmp_kv_type": [torch.bfloat16],
"B": [32],
"S1": [8],
"S2": [131072],
"T1": [256],
"T2": [1024],
"N1": [64],
"N2": [1],
"D": [512],
"K": [512],
"block_num1": [None],
"block_num2": [None],
"block_size1": [128],
"block_size2": [128],
"cu_seqlens_q": [None],
"seqused_ori_kv": [None],
"seqused_cmp_kv": [None],
"cmp_residual_kv": [None],
"softmax_scale": [0.04419417],
"cmp_ratio": [128],
"return_softmax_lse": [False],
"ori_mask_mode": [4],
"cmp_mask_mode": [3],
"ori_win_left": [127],
"ori_win_right": [0],
"template_mode": ["HCA"],
}
}
dev = torch.device('npu:0')
cu_seqlens_q = [0, 8, 16, 24, 32, 40, 48, 56, 64, 72, 80, 88, 96, 104, 112, 120, 128, 136, 144, 152, 160, 168, 176, 184, 192, 200, 208, 216, 224, 232, 240, 248, 256]
seqused_ori_kv = [131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072],
[ 0, 4, 8, 12, 16, 20, 24, 28, 32, 36, 40, 44, 48, 52,
56, 60, 64, 68, 72, 76, 80, 84, 88, 92, 96, 100, 104, 108,
112, 116, 120, 124, 128]
cu_seqlens_q_tensor = torch.tensor(cu_seqlens_q, dtype=torch.int32, device=dev)
seqused_ori_kv_tensor = torch.tensor(seqused_ori_kv, dtype=torch.int32, device=dev)
metadata = torch.ops.cann_ops_transformer.sparse_flash_mla_metadata(
num_heads_q=64,
num_heads_kv=1,
head_dim=512,
cu_seqlens_q=cu_seqlens_q_tensor,
cu_seqlens_ori_kv=None,
cu_seqlens_cmp_kv=None,
seqused_q=None,
seqused_ori_kv=seqused_ori_kv_tensor,
seqused_cmp_kv=None,
cmp_residual_kv=None,
ori_topk_length=None,
cmp_topk_length=None,
batch_size=32,
max_seqlen_q=8,
max_seqlen_ori_kv=131072,
max_seqlen_cmp_kv=32768,
ori_topk=0,
cmp_topk=512,
cmp_ratio=4,
ori_mask_mode=4,
cmp_mask_mode=3,
ori_win_left=127,
ori_win_right=0,
layout_q='TND',
layout_kv='PA_BBND',
has_ori_kv=True,
has_cmp_kv=True,
)
print_metadata(metadata)
{'testcase_name': 'HCA_decode_pa2',
'layout_q': 'TND',
'layout_kv': 'PA_BBND',
'q_type': torch.bfloat16,
'ori_kv_type': torch.bfloat16,
'cmp_kv_type': torch.bfloat16,
'B': 32,
'S1': 8,
'S2': 131072,
'T1': 256,
'T2': 1024,
'T3': None,
'N1': 64,
'N2': 1,
'D': 512,
'K1': None,
'K': 512,
'block_num1': 32768, 'block_num2': 256, 'block_size1': 128, 'block_size2': 128,
'seqused_q': None,
'cu_seqlens_q': [0, 8, 16, 24, 32, 40, 48, 56, 64, 72, 80, 88, 96, 104, 112, 120, 128, 136, 144, 152, 160, 168, 176, 184, 192, 200, 208, 216, 224, 232, 240, 248, 256],
'cu_seqlens_ori_kv': None,
'cu_seqlens_cmp_kv': None,
'seqused_ori_kv': [131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072, 131072],
'seqused_cmp_kv': None,
'cmp_residual_kv': None,
'softmax_scale': 0.04419417,
'cmp_ratio': 128,
'ori_mask_mode': 4,
'cmp_mask_mode': 3,
'ori_win_left': 127,
'ori_win_right': 0,
'ori_kv_topk_mode': 'no',
'cmp_kv_topk_mode': 'no',
'ori_sparse_indices_mode': 'full',
'cmp_sparse_indices_mode': 'full',
'actlen_mode': 'full',
'template_mode': 'HCA', 'q_datarange': [-2, 2], 'ori_kv_datarange': [-2, 2], 'cmp_kv_datarange': [-2, 2], 'random_seq': False, 'return_softmax_lse': False,
'ori_topk_length': None,
'cmp_topk_length': None}