flashinfer.testing.attention_tb_per_sec¶
- flashinfer.testing.attention_tb_per_sec(batch_size, qo_seqlen, kv_seqlen, head_dim_qk, head_dim_vo, num_qo_heads, num_kv_heads, time, q_dtype=torch.bfloat16, kv_dtype=torch.bfloat16, o_dtype=torch.bfloat16)¶
计算给定注意力层实现的每秒 TB 性能。假设批处理中所有序列长度相同。
- 参数:
batch_size (int) – 批大小。
qo_seqlen (int) – 查询的序列长度。
kv_seqlen (int) – 键和值的序列长度。
head_dim_qk (int) – 查询和键的头部维度。
head_dim_vo (int) – 值的头部维度。
num_qo_heads (int) – 查询头数。
num_kv_heads (int) – 键和值的头数。
time (float) – 毫秒级的执行时间。
q_dtype (torch.dtype) – 查询的数据类型。
kv_dtype (torch.dtype) – 键和值的数据类型。
o_dtype (torch.dtype) – 输出的数据类型。
- 返回值:
该层的每秒 TB 数。
- 返回值类型:
tb_per_sec (float)