已合并
modify document #4586
molly123321创建于 3月30日
modify document #4586
已合并
molly123321创建于 3月30日
380 个文件变更+2974-3192
Mdocs/zh/SECURITYNOTE.md+17-6
@@ -1,19 +1,22 @@
1# OpPlugin安全声明1# OpPlugin安全声明
2 2 
3## 系统安全加固3## 系统安全加固
4+ 
4建议用户在系统中配置开启ASLR(级别2 ),又称**全随机地址空间布局随机化**,可参考以下方式进行配置:5建议用户在系统中配置开启ASLR(级别2 ),又称**全随机地址空间布局随机化**,可参考以下方式进行配置:
5 6 
6 echo 2 > /proc/sys/kernel/randomize_va_space7 echo 2 > /proc/sys/kernel/randomize_va_space
7 8 
8## 运行用户建议9## 运行用户建议
10+ 
9OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓运行用户建议](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E8%BF%90%E8%A1%8C%E7%94%A8%E6%88%B7%E5%BB%BA%E8%AE%AE)。11OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓运行用户建议](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E8%BF%90%E8%A1%8C%E7%94%A8%E6%88%B7%E5%BB%BA%E8%AE%AE)。
10 12 
11## 文件权限控制13## 文件权限控制
12-1. 用户安装和使用过程需要做好权限控制,建议参考[文件(夹)各场景权限管控推荐最大值](#文件(夹)各场景权限管控推荐最大值)进行设置。如需要保存安装/卸载日志,可在安装/卸载命令后面加上参数--log <FILE>, 注意对<FILE>文件及目录做好权限管控。14+ 
15+1. 用户安装和使用过程需要做好权限控制,建议参考[文件(夹)各场景权限管控推荐最大值](#11111)进行设置。如需要保存安装/卸载日志,可在安装/卸载命令后面加上参数--log \<FILE>, 注意对\<FILE>文件及目录做好权限管控。
13 16 
142. 建议用户在主机(包括宿主机)及容器中设置运行系统umask值为0027及以上,保障新增文件夹默认最高权限为750,新增文件默认最高权限为640。172. 建议用户在主机(包括宿主机)及容器中设置运行系统umask值为0027及以上,保障新增文件夹默认最高权限为750,新增文件默认最高权限为640。
15 18 
16-##### 文件(夹)各场景权限管控推荐最大值19+<h3 id="11111">文件(夹)各场景权限管控推荐最大值</h3>
17 20 
18| 类型 | Linux权限参考最大值 |21| 类型 | Linux权限参考最大值 |
19|----------------------------------- |-----------------------|22|----------------------------------- |-----------------------|
@@ -36,33 +39,41 @@ OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓运行用
36| 加解密接口、加解密脚本 | 500(r-x------) |39| 加解密接口、加解密脚本 | 500(r-x------) |
37 40 
38## 调试工具声明41## 调试工具声明
42+ 
39OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓调试工具声明](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E8%B0%83%E8%AF%95%E5%B7%A5%E5%85%B7%E5%A3%B0%E6%98%8E)。43OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓调试工具声明](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E8%B0%83%E8%AF%95%E5%B7%A5%E5%85%B7%E5%A3%B0%E6%98%8E)。
40 44 
41## 数据安全声明45## 数据安全声明
46+ 
42OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓数据安全声明](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E6%95%B0%E6%8D%AE%E5%AE%89%E5%85%A8%E5%A3%B0%E6%98%8E)。47OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓数据安全声明](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E6%95%B0%E6%8D%AE%E5%AE%89%E5%85%A8%E5%A3%B0%E6%98%8E)。
43 48 
44## 构建安全声明49## 构建安全声明
50+ 
45OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓构建安全声明](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E6%9E%84%E5%BB%BA%E5%AE%89%E5%85%A8%E5%A3%B0%E6%98%8E)。51OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓构建安全声明](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E6%9E%84%E5%BB%BA%E5%AE%89%E5%85%A8%E5%A3%B0%E6%98%8E)。
46 52 
47## 运行安全声明53## 运行安全声明
54+ 
48OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓运行安全声明](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E8%BF%90%E8%A1%8C%E5%AE%89%E5%85%A8%E5%A3%B0%E6%98%8E)。55OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓运行安全声明](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E8%BF%90%E8%A1%8C%E5%AE%89%E5%85%A8%E5%A3%B0%E6%98%8E)。
49 56 
50## 公网地址声明57## 公网地址声明
58+ 
51在OpPlugin的配置文件和脚本中存在[公网地址](#公网地址)。59在OpPlugin的配置文件和脚本中存在[公网地址](#公网地址)。
52 60 
53-##### 公网地址61+### 公网地址
54 62 
55| 类型 | 开源代码地址 | 文件名 | 公网IP地址/公网URL地址/域名/邮箱地址 | 用途说明 |63| 类型 | 开源代码地址 | 文件名 | 公网IP地址/公网URL地址/域名/邮箱地址 | 用途说明 |
56|------------------------|-------------------------|-------------------------|------------------------------------------------------------------------------------------------------------|-------------------------|64|------------------------|-------------------------|-------------------------|------------------------------------------------------------------------------------------------------------|-------------------------|
57-| 开发引入 | 不涉及 | ci\build.sh | https://gitcode.com/ascend/pytorch.git | 编译脚本根据torch_npu仓库地址拉取代码进行编译 |65+| 开发引入 | 不涉及 | ci\build.sh | [https://gitcode.com/ascend/pytorch.git](https://gitcode.com/ascend/pytorch.git) | 编译脚本根据torch_npu仓库地址拉取代码进行编译 |
58-| 开发引入 | 不涉及 | ci\exec_ut.sh | https://gitcode.com/ascend/pytorch.git | UT脚本根据torch_npu仓库地址拉取代码进行UT测试 |66+| 开发引入 | 不涉及 | ci\exec_ut.sh | [https://gitcode.com/ascend/pytorch.git](https://gitcode.com/ascend/pytorch.git) | UT脚本根据torch_npu仓库地址拉取代码进行UT测试 |
59-| 开源代码引入 |pytorch\aten\src\ATen\native\TensorCompare.cpp | op_plugin\ops\opapi\IsInKernelNpuOpApi.cpp | https://github.com/numpy/numpy/blob/fb215c76967739268de71aa4bda55dd1b062bc2e/numpy/lib/arraysetops.py#L575 | 算法实现借鉴numpy的源码地址|67+| 开源代码引入 |pytorch\aten\src\ATen\native\TensorCompare.cpp | op_plugin\ops\opapi\IsInKernelNpuOpApi.cpp | [https://github.com/numpy/numpy/blob/fb215c76967739268de71aa4bda55dd1b062bc2e/numpy/lib/arraysetops.py#L575](https://github.com/numpy/numpy/blob/fb215c76967739268de71aa4bda55dd1b062bc2e/numpy/lib/arraysetops.py#L575) | 算法实现借鉴numpy的源码地址|
60 68 
61## 公开接口声明69## 公开接口声明
70+ 
62OpPlugin的运行依赖torch_npu,不提供公开接口。71OpPlugin的运行依赖torch_npu,不提供公开接口。
63 72 
64## 通信安全加固73## 通信安全加固
74+ 
65OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓通信安全加固](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E9%80%9A%E4%BF%A1%E5%AE%89%E5%85%A8%E5%8A%A0%E5%9B%BA)。75OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓通信安全加固](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E9%80%9A%E4%BF%A1%E5%AE%89%E5%85%A8%E5%8A%A0%E5%9B%BA)。
66 76 
67## 通信矩阵77## 通信矩阵
78+ 
68OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓通信矩阵](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E9%80%9A%E4%BF%A1%E7%9F%A9%E9%98%B5)。79OpPlugin的运行依赖torch_npu,本章内容请参考[torch_npu仓通信矩阵](https://gitcode.com/ascend/pytorch/blob/master/SECURITYNOTE.md#%E9%80%9A%E4%BF%A1%E7%9F%A9%E9%98%B5)。
Mdocs/zh/custom_APIs/Python_interface.md+0-1
@@ -1,2 +1 @@
1# Python接口1# Python接口
2- 
Mdocs/zh/custom_APIs/appendix.md+0-1
@@ -1,4 +1,3 @@
1# 附录1# 附录
2 2 
3- **[添加二进制黑名单示例](blacklist.md)** 3- **[添加二进制黑名单示例](blacklist.md)**
4- 
Mdocs/zh/custom_APIs/blacklist.md+0-1
@@ -23,4 +23,3 @@ option = {}
23option['NPU_FUZZY_COMPILE_BLACKLIST'] = "DynamicGRUV2,DynamicRNN" #根据实际场景进行替换23option['NPU_FUZZY_COMPILE_BLACKLIST'] = "DynamicGRUV2,DynamicRNN" #根据实际场景进行替换
24torch.npu.set_option(option)24torch.npu.set_option(option)
25```25```
26- 
Mdocs/zh/custom_APIs/cpp/C_interface.md+1-1
@@ -1 +1 @@
1-# C++ 接口1+# C++ 接口
Mdocs/zh/custom_APIs/cpp/C_list.md+10-11
@@ -164,56 +164,55 @@
164<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>申请一个device信息为NPU且实际内存在host侧的特殊Tensor。</p>164<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>申请一个device信息为NPU且实际内存在host侧的特殊Tensor。</p>
165</td>165</td>
166</tr>166</tr>
167-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard.md">c10_npu::NPUStreamGuard</a></p>167+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard.md">c10_npu::NPUStreamGuard</a></p>
168</td>168</td>
169<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>NPU设备流guard,保障作用域内的设备流,与`c10::cuda::CUDAStreamGuard`相同。</p>169<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>NPU设备流guard,保障作用域内的设备流,与`c10::cuda::CUDAStreamGuard`相同。</p>
170</td>170</td>
171</tr>171</tr>
172-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-current_device.md">c10_npu::NPUStreamGuard::current_device</a></p>172+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-current_device.md">c10_npu::NPUStreamGuard::current_device</a></p>
173</td>173</td>
174<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>返回guard当前设备。</p>174<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>返回guard当前设备。</p>
175</td>175</td>
176</tr>176</tr>
177-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-current_stream.md">c10_npu::NPUStreamGuard::current_stream</a></p>177+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-current_stream.md">c10_npu::NPUStreamGuard::current_stream</a></p>
178</td>178</td>
179<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>返回guard当前保障的流。</p>179<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>返回guard当前保障的流。</p>
180</td>180</td>
181</tr>181</tr>
182-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-NPUStreamGuard.md">c10_npu::NPUStreamGuard::NPUStreamGuard</a></p>182+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-NPUStreamGuard.md">c10_npu::NPUStreamGuard::NPUStreamGuard</a></p>
183</td>183</td>
184<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>构造函数,创建一个流guard。</p>184<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>构造函数,创建一个流guard。</p>
185</td>185</td>
186</tr>186</tr>
187-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-original_device.md">c10_npu::NPUStreamGuard::original_device</a></p>187+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-original_device.md">c10_npu::NPUStreamGuard::original_device</a></p>
188</td>188</td>
189<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>返回guard构造时的设备。</p>189<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>返回guard构造时的设备。</p>
190</td>190</td>
191</tr>191</tr>
192-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-original_stream.md">c10_npu::NPUStreamGuard::original_stream</a></p>192+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-original_stream.md">c10_npu::NPUStreamGuard::original_stream</a></p>
193</td>193</td>
194<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>返回guard构造时设置的流。</p>194<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>返回guard构造时设置的流。</p>
195</td>195</td>
196</tr>196</tr>
197-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-reset_stream.md">c10_npu::NPUStreamGuard::reset_stream</a></p>197+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-NPUStreamGuard-reset_stream.md">c10_npu::NPUStreamGuard::reset_stream</a></p>
198</td>198</td>
199<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>给guard重新设置新的流。</p>199<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>给guard重新设置新的流。</p>
200</td>200</td>
201</tr>201</tr>
202-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-stream_synchronize.md">c10_npu::stream_synchronize</a></p>202+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10_npu-stream_synchronize.md">c10_npu::stream_synchronize</a></p>
203</td>203</td>
204<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>NPU设备流同步,与`c10::cuda::stream_synchronize`相同。</p>204<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>NPU设备流同步,与`c10::cuda::stream_synchronize`相同。</p>
205</td>205</td>
206</tr>206</tr>
207-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10d_npu-ProcessGroupHCCL.md">c10d_npu::ProcessGroupHCCL</a></p>207+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10d_npu-ProcessGroupHCCL.md">c10d_npu::ProcessGroupHCCL</a></p>
208</td>208</td>
209<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>ProcessGroupHCCL继承自`c10d::Backend`,实现`HCCL`后端的相关接口,用于通信算子调用。</p>209<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>ProcessGroupHCCL继承自`c10d::Backend`,实现`HCCL`后端的相关接口,用于通信算子调用。</p>
210</td>210</td>
211</tr>211</tr>
212-</tr><tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10d_npu-ProcessGroupHCCL-batch_isend_irecv.md">c10d_npu::ProcessGroupHCCL::batch_isend_irecv</a></p>212+<tr><td class="cellrowborder" valign="top" width="38.61%" headers="mcps1.2.3.1.1 "><p><a href="c10d_npu-ProcessGroupHCCL-batch_isend_irecv.md">c10d_npu::ProcessGroupHCCL::batch_isend_irecv</a></p>
213</td>213</td>
214<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>发送或接收一批tensor,异步处理P2P操作序列中的每一个操作,并返回对应的请求。</p>214<td class="cellrowborder" valign="top" width="61.39%" headers="mcps1.2.3.1.2 "><p>发送或接收一批tensor,异步处理P2P操作序列中的每一个操作,并返回对应的请求。</p>
215</td>215</td>
216</tr>216</tr>
217</tbody>217</tbody>
218</table>218</table>
219- 
Mdocs/zh/custom_APIs/cpp/at_npu-native-empty_with_swapped_memory.md+3-4
@@ -1,4 +1,5 @@
1# at_npu::native::empty_with_swapped_memory1# at_npu::native::empty_with_swapped_memory
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -6,7 +7,6 @@
6|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
7|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
8 9 
9- 
10## 功能说明10## 功能说明
11 11 
12申请一个device信息为NPU且实际内存在host侧的特殊Tensor。12申请一个device信息为NPU且实际内存在host侧的特殊Tensor。
@@ -17,7 +17,7 @@ torch_npu\csrc\core\npu\NPUFormat.h
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21at::Tensor empty_with_swapped_memory(c10::IntArrayRef size, c10::optional<at::ScalarType> dtype_opt, c10::optional<c10::Device> device_opt)21at::Tensor empty_with_swapped_memory(c10::IntArrayRef size, c10::optional<at::ScalarType> dtype_opt, c10::optional<c10::Device> device_opt)
22```22```
23 23 
@@ -27,9 +27,8 @@ at::Tensor empty_with_swapped_memory(c10::IntArrayRef size, c10::optional<at::Sc
27- **dtype_opt** (`c10::optional<at::ScalarType>`):必选参数,表示生成Tensor的数据类型,若为`c10::nullopt`,则表示使用dtype全局默认值。27- **dtype_opt** (`c10::optional<at::ScalarType>`):必选参数,表示生成Tensor的数据类型,若为`c10::nullopt`,则表示使用dtype全局默认值。
28- **device_opt** (`c10::optional<c10::Device>`):必选参数,表示生成Tensor的设备信息,若为`c10::nullopt`,则表示使用当前默认device。28- **device_opt** (`c10::optional<c10::Device>`):必选参数,表示生成Tensor的设备信息,若为`c10::nullopt`,则表示使用当前默认device。
29 29 
30- 
31- 
32## 返回值说明30## 返回值说明
31+ 
33`at::Tensor`32`at::Tensor`
34 33 
35代表生成的特殊Tensor。34代表生成的特殊Tensor。
Mdocs/zh/custom_APIs/cpp/c10_npu-NPUStreamGuard-NPUStreamGuard.md+2-2
@@ -17,7 +17,7 @@ torch_npu\csrc\core\npu\NPUGuard.h
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21c10_npu::NPUStreamGuard::NPUStreamGuard(c10::Stream stream)21c10_npu::NPUStreamGuard::NPUStreamGuard(c10::Stream stream)
22```22```
23 23 
@@ -31,4 +31,4 @@ c10_npu::NPUStreamGuard::NPUStreamGuard(c10::Stream stream)
31 31 
32## 约束说明32## 约束说明
33 33 
34-34+
Mdocs/zh/custom_APIs/cpp/c10_npu-NPUStreamGuard-current_device.md+2-2
@@ -17,7 +17,7 @@ torch_npu\csrc\core\npu\NPUGuard.h
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21c10::Device c10_npu::NPUStreamGuard::current_device() const21c10::Device c10_npu::NPUStreamGuard::current_device() const
22```22```
23 23 
@@ -33,4 +33,4 @@ c10::Device c10_npu::NPUStreamGuard::current_device() const
33 33 
34## 约束说明34## 约束说明
35 35 
36-36+
Mdocs/zh/custom_APIs/cpp/c10_npu-NPUStreamGuard-current_stream.md+1-1
@@ -17,7 +17,7 @@ torch_npu\csrc\core\npu\NPUGuard.h
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21c10_npu::NPUStream c10_npu::NPUStreamGuard::current_stream() const21c10_npu::NPUStream c10_npu::NPUStreamGuard::current_stream() const
22```22```
23 23 
Mdocs/zh/custom_APIs/cpp/c10_npu-NPUStreamGuard-original_device.md+1-1
@@ -17,7 +17,7 @@ torch_npu\csrc\core\npu\NPUGuard.h
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21c10::Device c10_npu::NPUStreamGuard::original_device() const21c10::Device c10_npu::NPUStreamGuard::original_device() const
22```22```
23 23 
Mdocs/zh/custom_APIs/cpp/c10_npu-NPUStreamGuard-original_stream.md+2-2
@@ -17,7 +17,7 @@ torch_npu\csrc\core\npu\NPUGuard.h
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21c10_npu::NPUStream c10_npu::NPUStreamGuard::original_stream() const21c10_npu::NPUStream c10_npu::NPUStreamGuard::original_stream() const
22```22```
23 23 
@@ -33,4 +33,4 @@ c10_npu::NPUStream c10_npu::NPUStreamGuard::original_stream() const
33 33 
34## 约束说明34## 约束说明
35 35 
36-36+
Mdocs/zh/custom_APIs/cpp/c10_npu-NPUStreamGuard-reset_stream.md+2-2
@@ -17,7 +17,7 @@ torch_npu\csrc\core\npu\NPUGuard.h
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21void c10_npu::NPUStreamGuard::reset_stream(c10::Stream stream)21void c10_npu::NPUStreamGuard::reset_stream(c10::Stream stream)
22```22```
23 23 
@@ -31,4 +31,4 @@ void c10_npu::NPUStreamGuard::reset_stream(c10::Stream stream)
31 31 
32## 约束说明32## 约束说明
33 33 
34-'stream'必须是NPU流(即由NPU设备创建的c10::Stream),否则行为未定义。34+'stream'必须是NPU流(即由NPU设备创建的c10::Stream),否则行为未定义。
Mdocs/zh/custom_APIs/cpp/c10_npu-NPUStreamGuard.md+1-3
@@ -7,7 +7,6 @@
7|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9 9 
10- 
11## 功能说明10## 功能说明
12 11 
13NPU设备流guard,保障作用域内的设备流,与`c10::cuda::CUDAStreamGuard`相同。12NPU设备流guard,保障作用域内的设备流,与`c10::cuda::CUDAStreamGuard`相同。
@@ -18,7 +17,6 @@ torch_npu\csrc\core\npu\NPUGuard.h
18 17 
19## 函数原型18## 函数原型
20 19 
21-```20+```cpp
22struct c10_npu::NPUStreamGuard21struct c10_npu::NPUStreamGuard
23```22```
24- 
Mdocs/zh/custom_APIs/cpp/c10_npu-stream_synchronize.md+3-3
@@ -1,4 +1,5 @@
1# c10_npu::stream_synchronize1# c10_npu::stream_synchronize
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -6,7 +7,6 @@
6|<term> Atlas A3 训练系列产品</term> | √ |7|<term> Atlas A3 训练系列产品</term> | √ |
7|<term> Atlas A2 训练系列产品</term> | √ |8|<term> Atlas A2 训练系列产品</term> | √ |
8 9 
9- 
10## 功能说明10## 功能说明
11 11 
12NPU设备流同步,与`c10::cuda::stream_synchronize`相同。12NPU设备流同步,与`c10::cuda::stream_synchronize`相同。
@@ -17,7 +17,7 @@ torch_npu\csrc\core\npu\NPUFunctions.h
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21void stream_synchronize(aclrtStream stream)21void stream_synchronize(aclrtStream stream)
22```22```
23 23 
@@ -31,4 +31,4 @@ void stream_synchronize(aclrtStream stream)
31 31 
32## 约束说明32## 约束说明
33 33 
34-34+
Mdocs/zh/custom_APIs/cpp/c10d_npu-ProcessGroupHCCL-batch_isend_irecv.md+2-2
@@ -17,7 +17,7 @@ torch_npu\csrc\distributed\ProcessGroupHCCL.hpp
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21c10::intrusive_ptr<c10d::Work> batch_isend_irecv(std::vector<std::string>& op_type, std::vector<at::Tensor>& tensors, std::vector<uint32_t> remote_rank_list)21c10::intrusive_ptr<c10d::Work> batch_isend_irecv(std::vector<std::string>& op_type, std::vector<at::Tensor>& tensors, std::vector<uint32_t> remote_rank_list)
22```22```
23 23 
@@ -35,4 +35,4 @@ c10::intrusive_ptr<c10d::Work> batch_isend_irecv(std::vector<std::string>& op_ty
35 35 
36## 约束说明36## 约束说明
37 37 
38-38+
Mdocs/zh/custom_APIs/cpp/c10d_npu-ProcessGroupHCCL.md+1-3
@@ -17,7 +17,7 @@ torch_npu\csrc\distributed\ProcessGroupHCCL.hpp
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```cpp
21class c10d_npu::ProcessGroupHCCL21class c10d_npu::ProcessGroupHCCL
22```22```
23 23 
@@ -42,5 +42,3 @@ recv<br>
42recv_anysource<br>42recv_anysource<br>
43alltoall_base<br>43alltoall_base<br>
44alltoall<br>44alltoall<br>
45- 
46- 
Mdocs/zh/custom_APIs/cpp/(beta)at-Device.md+2-3
@@ -2,7 +2,7 @@
2 2 
3## 函数原型3## 函数原型
4 4 
5-```5+```cpp
6at::Device(const std::string &device_string)6at::Device(const std::string &device_string)
7```7```
8 8 
@@ -12,7 +12,7 @@ at::Device(const std::string &device_string)
12 12 
13## 参数说明13## 参数说明
14 14 
15-device_ring:string类型,提供的字符串必须遵循以下架构:(npu)[:<device-index\>],其中NPU指定设备类型,<device-index\>可选,指定设备索引。15+device_ring:string类型,提供的字符串必须遵循以下架构:`(npu)[:<device-index>]`,其中NPU指定设备类型,<device-index\>可选,指定设备索引。
16 16 
17## 支持的型号17## 支持的型号
18 18 
@@ -20,4 +20,3 @@ device_ring:string类型,提供的字符串必须遵循以下架构:(npu)[
20- <term>Atlas A2 训练系列产品</term>20- <term>Atlas A2 训练系列产品</term>
21- <term>Atlas A3 训练系列产品</term>21- <term>Atlas A3 训练系列产品</term>
22- <term>Atlas 推理系列产品</term>22- <term>Atlas 推理系列产品</term>
23- 
Mdocs/zh/custom_APIs/cpp/(beta)at_npu-detail-createNPUGenerator.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\aten\NPUGeneratorImpl.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10at::Generator at_npu::detail::createNPUGenerator(c10::DeviceIndex device_index = -1)10at::Generator at_npu::detail::createNPUGenerator(c10::DeviceIndex device_index = -1)
11```11```
12 12 
@@ -24,4 +24,3 @@ device_index:DeviceIndex类型,指定创建生成器的NPU设备id。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)at_npu-detail-getDefaultNPUGenerator.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\aten\NPUGeneratorImpl.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10at::Generator& at_npu::detail::getDefaultNPUGenerator(c10::DeviceIndex device_index = -1)10at::Generator& at_npu::detail::getDefaultNPUGenerator(c10::DeviceIndex device_index = -1)
11```11```
12 12 
@@ -24,4 +24,3 @@ device_index:DeviceIndex类型,指定获取生成器的NPU设备id。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)at_npu-native-empty_with_format.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFormat.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10at::Tensor at_npu::native::empty_with_format(c10::IntArrayRef sizes, const c10::TensorOptions& options, int64_t acl_format, bool keep_format = false)10at::Tensor at_npu::native::empty_with_format(c10::IntArrayRef sizes, const c10::TensorOptions& options, int64_t acl_format, bool keep_format = false)
11```11```
12 12 
@@ -30,4 +30,3 @@ keep_format:bool类型,是否指定格式,true表示指定获取tensor的
30- <term>Atlas A2 训练系列产品</term>30- <term>Atlas A2 训练系列产品</term>
31- <term>Atlas A3 训练系列产品</term>31- <term>Atlas A3 训练系列产品</term>
32- <term>Atlas 推理系列产品</term>32- <term>Atlas 推理系列产品</term>
33- 
Mdocs/zh/custom_APIs/cpp/(beta)at_npu-native-get_npu_format.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFormat.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10int64_t at_npu::native::get_npu_format(const at::Tensor& self)10int64_t at_npu::native::get_npu_format(const at::Tensor& self)
11```11```
12 12 
@@ -27,4 +27,3 @@ self:Tensor类型,待获取格式信息的tensor。
27- <term>Atlas A2 训练系列产品</term>27- <term>Atlas A2 训练系列产品</term>
28- <term>Atlas A3 训练系列产品</term>28- <term>Atlas A3 训练系列产品</term>
29- <term>Atlas 推理系列产品</term>29- <term>Atlas 推理系列产品</term>
30- 
Mdocs/zh/custom_APIs/cpp/(beta)at_npu-native-get_npu_storage_sizes.md+2-3
@@ -6,13 +6,13 @@ torch_npu\csrc\core\npu\NPUFormat.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10std::vector<int64_t> at_npu::native::get_npu_storage_sizes(const at::Tensor& self)10std::vector<int64_t> at_npu::native::get_npu_storage_sizes(const at::Tensor& self)
11```11```
12 12 
13## 功能说明13## 功能说明
14 14 
15-获取NPU tensor的内存大小,返回值类型vector<int64_t>,表示获取的NPU tensor内存大小。15+获取NPU tensor的内存大小,返回值类型vector\<int64_t>,表示获取的NPU tensor内存大小。
16 16 
17## 参数说明17## 参数说明
18 18 
@@ -24,4 +24,3 @@ self:Tensor类型,待获取内存大小的tensor。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)at_npu-native-npu_dropout_gen_mask.md+1-2
@@ -6,7 +6,7 @@ third_party\op-plugin\op_plugin\include\ops.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10at::Tensor npu_dropout_gen_mask(const at::Tensor &self, at::IntArrayRef size, double p, int64_t seed, int64_t offset, c10::optional<bool> parallel, c10::optional<bool> sync)10at::Tensor npu_dropout_gen_mask(const at::Tensor &self, at::IntArrayRef size, double p, int64_t seed, int64_t offset, c10::optional<bool> parallel, c10::optional<bool> sync)
11```11```
12 12 
@@ -30,4 +30,3 @@ at::Tensor npu_dropout_gen_mask(const at::Tensor &self, at::IntArrayRef size, do
30- <term>Atlas A2 训练系列产品</term>30- <term>Atlas A2 训练系列产品</term>
31- <term>Atlas A3 训练系列产品</term>31- <term>Atlas A3 训练系列产品</term>
32- <term>Atlas 推理系列产品</term>32- <term>Atlas 推理系列产品</term>
33- 
Mdocs/zh/custom_APIs/cpp/(beta)at_npu-native-npu_format_cast.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFormat.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10at::Tensor at_npu::native::npu_format_cast(const at::Tensor& self, int64_t acl_format)10at::Tensor at_npu::native::npu_format_cast(const at::Tensor& self, int64_t acl_format)
11```11```
12 12 
@@ -26,4 +26,3 @@ acl_format:int64_t型,待转换的格式。
26- <term>Atlas A2 训练系列产品</term>26- <term>Atlas A2 训练系列产品</term>
27- <term>Atlas A3 训练系列产品</term>27- <term>Atlas A3 训练系列产品</term>
28- <term>Atlas 推理系列产品</term>28- <term>Atlas 推理系列产品</term>
29- 
Mdocs/zh/custom_APIs/cpp/(beta)c10-npu-current_device.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\libs\init_npu.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10c10::DeviceIndex c10::npu::current_device()10c10::DeviceIndex c10::npu::current_device()
11```11```
12 12 
@@ -20,4 +20,3 @@ c10::DeviceIndex c10::npu::current_device()
20- <term>Atlas A2 训练系列产品</term>20- <term>Atlas A2 训练系列产品</term>
21- <term>Atlas A3 训练系列产品</term>21- <term>Atlas A3 训练系列产品</term>
22- <term>Atlas 推理系列产品</term>22- <term>Atlas 推理系列产品</term>
23- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-GetDevice.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFunctions.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10aclError c10_npu::GetDevice(c10::DeviceIndex* device)10aclError c10_npu::GetDevice(c10::DeviceIndex* device)
11```11```
12 12 
@@ -24,4 +24,3 @@ device:DeviceIndex类型,存储获取的设备id。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-SetDevice.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFunctions.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10aclError c10_npu::SetDevice(c10::DeviceIndex device)10aclError c10_npu::SetDevice(c10::DeviceIndex device)
11```11```
12 12 
@@ -24,4 +24,3 @@ device:DeviceIndex类型,待设置的NPU设备id。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-c10_npu_get_error_message.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUException.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10const char* c10_npu::c10_npu_get_error_message()10const char* c10_npu::c10_npu_get_error_message()
11```11```
12 12 
@@ -20,4 +20,3 @@ const char* c10_npu::c10_npu_get_error_message()
20- <term>Atlas A2 训练系列产品</term>20- <term>Atlas A2 训练系列产品</term>
21- <term>Atlas A3 训练系列产品</term>21- <term>Atlas A3 训练系列产品</term>
22- <term>Atlas 推理系列产品</term>22- <term>Atlas 推理系列产品</term>
23- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-current_device.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFunctions.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10c10::DeviceIndex c10_npu::current_device()10c10::DeviceIndex c10_npu::current_device()
11```11```
12 12 
@@ -20,4 +20,3 @@ NPU设备id获取,返回值类型DeviceIndex,表示获取到的设备id,
20- <term>Atlas A2 训练系列产品</term>20- <term>Atlas A2 训练系列产品</term>
21- <term>Atlas A3 训练系列产品</term>21- <term>Atlas A3 训练系列产品</term>
22- <term>Atlas 推理系列产品</term>22- <term>Atlas 推理系列产品</term>
23- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-device_count.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFunctions.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10c10::DeviceIndex c10_npu::device_count()10c10::DeviceIndex c10_npu::device_count()
11```11```
12 12 
@@ -20,4 +20,3 @@ c10::DeviceIndex c10_npu::device_count()
20- <term>Atlas A2 训练系列产品</term>20- <term>Atlas A2 训练系列产品</term>
21- <term>Atlas A3 训练系列产品</term>21- <term>Atlas A3 训练系列产品</term>
22- <term>Atlas 推理系列产品</term>22- <term>Atlas 推理系列产品</term>
23- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-getCurrentNPUStream.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUStream.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10c10_npu::NPUStream c10_npu::getCurrentNPUStream(c10::DeviceIndex device_index = -1)10c10_npu::NPUStream c10_npu::getCurrentNPUStream(c10::DeviceIndex device_index = -1)
11```11```
12 12 
@@ -24,4 +24,3 @@ device_index:DeviceIndex类型,获取流的NPU设备id。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-getDefaultNPUStream.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUStream.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10c10_npu::NPUStream c10_npu::getDefaultNPUStream(c10::DeviceIndex device_index = -1)10c10_npu::NPUStream c10_npu::getDefaultNPUStream(c10::DeviceIndex device_index = -1)
11```11```
12 12 
@@ -24,4 +24,3 @@ device_index:DeviceIndex类型,获取流的NPU设备id。默认值为-1,
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-getNPUStreamFromPool.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUStream.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10c10_npu::NPUStream c10_npu::getNPUStreamFromPool(c10::DeviceIndex device = -1)10c10_npu::NPUStream c10_npu::getNPUStreamFromPool(c10::DeviceIndex device = -1)
11```11```
12 12 
@@ -24,4 +24,3 @@ device:DeviceIndex类型,获取流的NPU设备id。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-setCurrentNPUStream.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUStream.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10void c10_npu::setCurrentNPUStream(c10_npu::NPUStream stream)10void c10_npu::setCurrentNPUStream(c10_npu::NPUStream stream)
11```11```
12 12 
@@ -24,4 +24,3 @@ stream:NPUStream类型,待设置的NPU流。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-set_device.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFunctions.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10void c10_npu::set_device(c10::DeviceIndex device)10void c10_npu::set_device(c10::DeviceIndex device)
11```11```
12 12 
@@ -24,4 +24,3 @@ device:DeviceIndex类型,待设置的NPU设备id。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-warn_or_error_on_sync.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFunctions.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10void c10_npu::warn_or_error_on_sync()10void c10_npu::warn_or_error_on_sync()
11```11```
12 12 
@@ -20,4 +20,3 @@ NPU同步时警告,无返回值,根据当前警告等级进行报错或警
20- <term>Atlas A2 训练系列产品</term>20- <term>Atlas A2 训练系列产品</term>
21- <term>Atlas A3 训练系列产品</term>21- <term>Atlas A3 训练系列产品</term>
22- <term>Atlas 推理系列产品</term>22- <term>Atlas 推理系列产品</term>
23- 
Mdocs/zh/custom_APIs/cpp/(beta)c10_npu-warning_state.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\core\npu\NPUFunctions.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10c10_npu::WarningState& c10_npu::warning_state()10c10_npu::WarningState& c10_npu::warning_state()
11```11```
12 12 
@@ -20,4 +20,3 @@ c10_npu::WarningState& c10_npu::warning_state()
20- <term>Atlas A2 训练系列产品</term>20- <term>Atlas A2 训练系列产品</term>
21- <term>Atlas A3 训练系列产品</term>21- <term>Atlas A3 训练系列产品</term>
22- <term>Atlas 推理系列产品</term>22- <term>Atlas 推理系列产品</term>
23- 
Mdocs/zh/custom_APIs/cpp/(beta)class-at_npu-NPUGeneratorImpl.md+6-7
@@ -16,9 +16,9 @@ NPUGeneratorImpl是一个随机数生成器类,实现了NPU设备随机数的
16 16 
17 device_index:DeviceIndex类型,指定npu设备id。17 device_index:DeviceIndex类型,指定npu设备id。
18 18 
19-- **std::shared_ptr<NPUGeneratorImpl> at_npu::NPUGeneratorImpl::clone()**19+- **std::shared_ptr\<NPUGeneratorImpl> at_npu::NPUGeneratorImpl::clone()**
20 20 
21- NPUGeneratorImpl拷贝函数,返回值类型shared_ptr<NPUGeneratorImpl>,返回NPUGeneratorImpl拷贝,与std::shared_ptr<CUDAGeneratorImpl> at::CUDAGeneratorImpl::clone()相同。21+ NPUGeneratorImpl拷贝函数,返回值类型shared_ptr\<NPUGeneratorImpl>,返回NPUGeneratorImpl拷贝,与std::shared_ptr\<CUDAGeneratorImpl> at::CUDAGeneratorImpl::clone()相同。
22 22 
23- **void at_npu::NPUGeneratorImpl::set_current_seed(uint64_t seed)**23- **void at_npu::NPUGeneratorImpl::set_current_seed(uint64_t seed)**
24 24 
@@ -50,9 +50,9 @@ NPUGeneratorImpl是一个随机数生成器类,实现了NPU设备随机数的
50 50 
51 new_state:TensorImpl类型,待设置的状态,需要通过at::detail::check_rng_state检测。51 new_state:TensorImpl类型,待设置的状态,需要通过at::detail::check_rng_state检测。
52 52 
53-- **c10::intrusive_ptr<c10::TensorImpl> c10::TensorImpl at_npu::NPUGeneratorImpl::get_state()**53+- **c10::intrusive_ptr\<c10::TensorImpl> c10::TensorImpl at_npu::NPUGeneratorImpl::get_state()**
54 54 
55- NPUGeneratorImpl状态获取,返回值类型intrusive_ptr<c10::TensorImpl>,返回生成器状态,与c10::intrusive_ptr<c10::TensorImpl> at::CUDAGeneratorImpl::get_state()相同。55+ NPUGeneratorImpl状态获取,返回值类型intrusive_ptr\<c10::TensorImpl>,返回生成器状态,与c10::intrusive_ptr\<c10::TensorImpl> at::CUDAGeneratorImpl::get_state()相同。
56 56 
57- **void at_npu::NPUGeneratorImpl::set_philox_offset_per_thread(uint64_t offset)**57- **void at_npu::NPUGeneratorImpl::set_philox_offset_per_thread(uint64_t offset)**
58 58 
@@ -97,9 +97,9 @@ Pytorch2.5.1及以上版本,新增以下成员函数:
97 在capture状态下为aclgraph设置期望的随机数生成状态,与at::CUDAGeneratorImpl::graphsafe_set_state(const c10::intrusive_ptr& state)功能相同。97 在capture状态下为aclgraph设置期望的随机数生成状态,与at::CUDAGeneratorImpl::graphsafe_set_state(const c10::intrusive_ptr& state)功能相同。
98 98
99 state:随机数生成器状态。99 state:随机数生成器状态。
100-- **c10::intrusive_ptr<c10::GeneratorImpl> graphsafe_get_state()**100+- **c10::intrusive_ptr\<c10::GeneratorImpl> graphsafe_get_state()**
101 101 
102- 在capture状态下为aclgraph查询随机数生成对象,与c10::intrusive_ptr<c10::GeneratorImpl> at::CUDAGeneratorImpl::graphsafe_get_state()功能相同。102+ 在capture状态下为aclgraph查询随机数生成对象,与c10::intrusive_ptr\<c10::GeneratorImpl> at::CUDAGeneratorImpl::graphsafe_get_state()功能相同。
103 103
104 返回值为c10::GeneratorImpl对象。104 返回值为c10::GeneratorImpl对象。
105- **void register_graph(c10_npu::NPUGraph* graph)**105- **void register_graph(c10_npu::NPUGraph* graph)**
@@ -115,4 +115,3 @@ Pytorch2.5.1及以上版本,新增以下成员函数:
115- <term>Atlas A2 训练系列产品</term>115- <term>Atlas A2 训练系列产品</term>
116- <term>Atlas A3 训练系列产品</term>116- <term>Atlas A3 训练系列产品</term>
117- <term>Atlas 推理系列产品</term>117- <term>Atlas 推理系列产品</term>
118- 
Mdocs/zh/custom_APIs/cpp/(beta)class-at_npu-native-OpCommand.md+0-2
@@ -188,5 +188,3 @@ OpCommand是一个封装下层算子调用的类,实现了NPU设备下层算
188- <term>Atlas A2 训练系列产品</term>188- <term>Atlas A2 训练系列产品</term>
189- <term>Atlas A3 训练系列产品</term>189- <term>Atlas A3 训练系列产品</term>
190- <term>Atlas 推理系列产品</term>190- <term>Atlas 推理系列产品</term>
191- 
192- 
Mdocs/zh/custom_APIs/cpp/(beta)class-c10_npu-NPUStream.md+0-1
@@ -114,4 +114,3 @@ NPUStream是一个NPU流类,实现了NPU流管理的相关功能,是属于NP
114- <term>Atlas A2 训练系列产品</term>114- <term>Atlas A2 训练系列产品</term>
115- <term>Atlas A3 训练系列产品</term>115- <term>Atlas A3 训练系列产品</term>
116- <term>Atlas 推理系列产品</term>116- <term>Atlas 推理系列产品</term>
117- 
Mdocs/zh/custom_APIs/cpp/(beta)struct-c10_npu-NPUEvent.md+0-1
@@ -92,4 +92,3 @@ NPUEvent是一个事件类,实现了NPU设备事件管理的相关功能,可
92- <term>Atlas A2 训练系列产品</term>92- <term>Atlas A2 训练系列产品</term>
93- <term>Atlas A3 训练系列产品</term>93- <term>Atlas A3 训练系列产品</term>
94- <term>Atlas 推理系列产品</term>94- <term>Atlas 推理系列产品</term>
95- 
Mdocs/zh/custom_APIs/cpp/(beta)struct-c10_npu-NPUHooksArgs.md+0-1
@@ -14,4 +14,3 @@ NPUHooksArgs是一个Hook参数类,提供了NPU Hook的相关参数。
14- <term>Atlas A2 训练系列产品</term>14- <term>Atlas A2 训练系列产品</term>
15- <term>Atlas A3 训练系列产品</term>15- <term>Atlas A3 训练系列产品</term>
16- <term>Atlas 推理系列产品</term>16- <term>Atlas 推理系列产品</term>
17- 
Mdocs/zh/custom_APIs/cpp/(beta)struct-c10_npu-NPUHooksInterface.md+0-1
@@ -24,4 +24,3 @@ device_index:DeviceIndex类型,指定NPU设备id。
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)torch-npu-synchronize.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\libs\init_npu.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10void torch::npu::synchronize(int64_t device_index = -1)10void torch::npu::synchronize(int64_t device_index = -1)
11```11```
12 12 
@@ -24,4 +24,3 @@ device_index:int64_t类型,用来同步设备的index,默认-1,即同步
24- <term>Atlas A2 训练系列产品</term>24- <term>Atlas A2 训练系列产品</term>
25- <term>Atlas A3 训练系列产品</term>25- <term>Atlas A3 训练系列产品</term>
26- <term>Atlas 推理系列产品</term>26- <term>Atlas 推理系列产品</term>
27- 
Mdocs/zh/custom_APIs/cpp/(beta)torch_npu-finalize_npu.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\libs\init_npu.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10void torch_npu::finalize_npu()10void torch_npu::finalize_npu()
11```11```
12 12 
@@ -20,4 +20,3 @@ void torch_npu::finalize_npu()
20- <term>Atlas A2 训练系列产品</term>20- <term>Atlas A2 训练系列产品</term>
21- <term>Atlas A3 训练系列产品</term>21- <term>Atlas A3 训练系列产品</term>
22- <term>Atlas 推理系列产品</term>22- <term>Atlas 推理系列产品</term>
23- 
Mdocs/zh/custom_APIs/cpp/(beta)torch_npu-init_npu.md+1-2
@@ -6,7 +6,7 @@ torch_npu\csrc\libs\init_npu.h
6 6 
7## 函数原型7## 函数原型
8 8 
9-```9+```cpp
10void torch_npu::init_npu(const c10::DeviceIndex device_index = 0)10void torch_npu::init_npu(const c10::DeviceIndex device_index = 0)
11void torch_npu::init_npu(const std::string& device_str)11void torch_npu::init_npu(const std::string& device_str)
12void torch_npu::init_npu(const at::Device& device)12void torch_npu::init_npu(const at::Device& device)
@@ -28,4 +28,3 @@ void torch_npu::init_npu(const at::Device& device)
28- <term>Atlas A2 训练系列产品</term>28- <term>Atlas A2 训练系列产品</term>
29- <term>Atlas A3 训练系列产品</term>29- <term>Atlas A3 训练系列产品</term>
30- <term>Atlas 推理系列产品</term>30- <term>Atlas 推理系列产品</term>
31- 
Mdocs/zh/custom_APIs/determin_API_list.md+14-15
@@ -5,12 +5,13 @@
5在使用PyTorch框架进行训练时,若需要输出结果排除随机性,则需要设置确定性计算开关。在开启确定性计算时,当使用相同的输入在相同的硬件和软件上执行相同的操作,输出的结果每次都是相同的。5在使用PyTorch框架进行训练时,若需要输出结果排除随机性,则需要设置确定性计算开关。在开启确定性计算时,当使用相同的输入在相同的硬件和软件上执行相同的操作,输出的结果每次都是相同的。
6 6 
7> [!NOTE] 7> [!NOTE]
8->- 确定性计算固定方法都必须与待固定的网络、算子等在同一个主进程,部分模型脚本中main()与训练网络并不在一个进程中。8+>
9->- 当前同一线程中只能设置一次确定性状态,多次设置以第次有效设置为准后续设置会生效9+> - 确定性计算固定方法都必须与待固定的网络、算子等在同个主进程部分模型脚本中main()与训练网络并在一个进程中
10-> 有效设置:在设置确定性状态真正执行了一算子的任务下发,如果仅设置,没算子下发,只能是确定性变量开启,并未下发给算子,因不执行算子,不知道哪个算子需要执行确定性10+> - 当前同一线程中只能设置一次确定性状态,次设置以第一次效设置后续设置会生效<br>
11+> 有效设置:在设置确定性状态后,真正执行了一次算子的任务下发,如果仅设置,没有算子下发,只能是确定性变量开启,并未下发给算子,因为不执行算子,不知道哪个算子需要执行确定性。<br>
11> 解决方案:12> 解决方案:
12-> 1. 暂不推荐一个线程多次设置确定性。13+> 1. 暂不推荐一个线程多次设置确定性。
13-> 2. 该问题在二进制开启和关闭情况下均存在,在后续版本中会解决该问题。14+> 2. 该问题在二进制开启和关闭情况下均存在,在后续版本中会解决该问题。
14 15 
15## 使用方法16## 使用方法
16 17 
@@ -19,23 +20,23 @@
19> [!CAUTION] 20> [!CAUTION]
20> 开启确定性开关可能会导致性能下降。21> 开启确定性开关可能会导致性能下降。
21 22 
22-1. 开启确定性计算开关:23+1. 开启确定性计算开关:
23 24 
24- ```25+ ```python
25 torch.use_deterministic_algorithms(True)26 torch.use_deterministic_algorithms(True)
26 ```27 ```
27 28 
28-2. 验证设置是否成功。29+2. 验证设置是否成功。
29 30 
30- 1. 执行如下命令查询接口设置:31+ 1. 执行如下命令查询接口设置:
31 32 
32- ```33+ ```python
33 torch.are_deterministic_algorithms_enabled()34 torch.are_deterministic_algorithms_enabled()
34 ```35 ```
35 36 
36- 2. 返回显示如下:37+ 2. 返回显示如下:
37 38 
38- ```39+ ```python
39 print(torch.are_deterministic_algorithms_enabled())40 print(torch.are_deterministic_algorithms_enabled())
40 ```41 ```
41 42 
@@ -43,6 +44,4 @@
43 44 
44## API支持清单45## API支持清单
45 46 
46-目前昇腾支持确定性计算的自定义API为[(beta)torch_npu.npu_group_norm_swish](torch_npu-npu_group_norm_swish.md)。47+目前昇腾支持确定性计算的自定义API为[(beta)torch_npu.npu_group_norm_swish](torch_npu/torch_npu-npu_group_norm_swish.md)。
47- 
48- 
Mdocs/zh/custom_APIs/distributed/Distributed.md+1-1
@@ -1 +1 @@
1-# Distributed1+# Distributed
Mdocs/zh/custom_APIs/distributed/Distributed_list.md+0-1
@@ -43,4 +43,3 @@
43</tr>43</tr>
44</tbody>44</tbody>
45</table>45</table>
46- 
Mdocs/zh/custom_APIs/distributed/torch-distributed-distributed_c10d.md+3-3
@@ -1,4 +1,5 @@
1# torch.distributed.distributed_c10d._world.default_pg._get_backend(torch.device("npu")).get_hccl_comm_name1# torch.distributed.distributed_c10d._world.default_pg._get_backend(torch.device("npu")).get_hccl_comm_name
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -14,7 +15,7 @@
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch.distributed.distributed_c10d._world.default_pg._get_backend(torch.device("npu")).get_hccl_comm_name(rankid->int,init_comm=True) -> string19torch.distributed.distributed_c10d._world.default_pg._get_backend(torch.device("npu")).get_hccl_comm_name(rankid->int,init_comm=True) -> string
19```20```
20 21 
@@ -33,6 +34,7 @@ torch.distributed.distributed_c10d._world.default_pg._get_backend(torch.device("
33>hccl初始化会申请内存资源,造成内存升高,默认申请内存大小为Send buffer与Recv buffer各200M,共400M。buffer大小受环境变量HCCL_BUFFSIZE控制。34>hccl初始化会申请内存资源,造成内存升高,默认申请内存大小为Send buffer与Recv buffer各200M,共400M。buffer大小受环境变量HCCL_BUFFSIZE控制。
34 35 
35## 返回值说明36## 返回值说明
37+ 
36`string`38`string`
37 39 
38代表string类型的集合通信域的名字。40代表string类型的集合通信域的名字。
@@ -42,7 +44,6 @@ torch.distributed.distributed_c10d._world.default_pg._get_backend(torch.device("
42- 使用该接口前确保`init_process_group`已被调用,且初始化的backend为hccl。44- 使用该接口前确保`init_process_group`已被调用,且初始化的backend为hccl。
43- PyTorch 2.1.0及以后版本与PyTorch 2.1.0之前的版本对该接口调用方式不同,见[调用示例](#section14459801435)。45- PyTorch 2.1.0及以后版本与PyTorch 2.1.0之前的版本对该接口调用方式不同,见[调用示例](#section14459801435)。
44 46 
45- 
46## 调用示例<a name="section14459801435"></a>47## 调用示例<a name="section14459801435"></a>
47 48 
48```python49```python
@@ -75,4 +76,3 @@ if __name__ == "__main__":
75group_name_076group_name_0
76group_name_077group_name_0
77```78```
78- 
Mdocs/zh/custom_APIs/distributed/torch_npu-distributed-reduce_scatter_tensor_uneven.md+3-7
@@ -1,4 +1,5 @@
1# (beta)torch_npu.distributed.reduce_scatter_tensor_uneven1# (beta)torch_npu.distributed.reduce_scatter_tensor_uneven
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -12,11 +13,10 @@
12 13 
13## 函数原型14## 函数原型
14 15 
15-```16+```python
16torch_npu.distributed.reduce_scatter_tensor_uneven(output, input, input_split_sizes =None, op=dist.ReduceOp.SUM, group=None, async_op=False) -> torch.distributed.distributed_c10d.Work17torch_npu.distributed.reduce_scatter_tensor_uneven(output, input, input_split_sizes =None, op=dist.ReduceOp.SUM, group=None, async_op=False) -> torch.distributed.distributed_c10d.Work
17```18```
18 19 
19- 
20## 参数说明20## 参数说明
21 21 
22- **output** (`Tensor`):必选参数,输出Tensor,用于接收计算数据。22- **output** (`Tensor`):必选参数,输出Tensor,用于接收计算数据。
@@ -28,20 +28,17 @@ torch_npu.distributed.reduce_scatter_tensor_uneven(output, input, input_split_si
28- **group** (`torch.distributed.distributed_c10d.ProcessGroup`):可选参数,分布式进程组,默认值None。28- **group** (`torch.distributed.distributed_c10d.ProcessGroup`):可选参数,分布式进程组,默认值None。
29- **async_op** (`bool`):可选参数,是否异步调用,默认值False。29- **async_op** (`bool`):可选参数,是否异步调用,默认值False。
30 30 
31- 
32## 返回值说明31## 返回值说明
33 32 
34该函数直接返回进行计算时的工作句柄,实际计算结果传给output。33该函数直接返回进行计算时的工作句柄,实际计算结果传给output。
35`output`:类型为Tensor,其shape无特殊约束。34`output`:类型为Tensor,其shape无特殊约束。
36 35 
37- 
38## 约束说明36## 约束说明
39 37 
40- 此接口仅可在单机场景下使用。38- 此接口仅可在单机场景下使用。
41 39 
42- `input_split_sizes`元素之和等于`input`的0维;`input_split_sizes`元素个数等于`group`的size。40- `input_split_sizes`元素之和等于`input`的0维;`input_split_sizes`元素个数等于`group`的size。
43 41 
44- 
45## 调用示例42## 调用示例
46 43 
47创建以下文件test.py并保存。44创建以下文件test.py并保存。
@@ -67,7 +64,6 @@ torch_npu.distributed.reduce_scatter_tensor_uneven(
67 64 
68执行如下命令。65执行如下命令。
69 66 
70-```67+```bash
71torchrun --nproc-per-node=2 test.py68torchrun --nproc-per-node=2 test.py
72```69```
73- 
Mdocs/zh/custom_APIs/distributed/(beta)torch-distributed-ProcessGroupHCCL.md+5-2
@@ -1,4 +1,5 @@
1# (beta)torch.distributed.ProcessGroupHCCL1# (beta)torch.distributed.ProcessGroupHCCL
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -13,7 +14,8 @@
13创建一个ProcessGroupHCCL对象并返回。14创建一个ProcessGroupHCCL对象并返回。
14 15 
15## 函数原型16## 函数原型
16-```17+ 
18+```python
17torch.distributed.ProcessGroupHCCL(store, rank, size, timeout) -> ProcessGroup19torch.distributed.ProcessGroupHCCL(store, rank, size, timeout) -> ProcessGroup
18```20```
19 21 
@@ -25,4 +27,5 @@ torch.distributed.ProcessGroupHCCL(store, rank, size, timeout) -> ProcessGroup
25- **timeout**:通讯中断时间,判断节点断连,默认值为1800s。27- **timeout**:通讯中断时间,判断节点断连,默认值为1800s。
26 28 
27## 返回值说明29## 返回值说明
28-`ProcessGroup`30+ 
31+`ProcessGroup`
Mdocs/zh/custom_APIs/distributed/(beta)torch-distributed-is_hccl_available.md+5-2
@@ -1,4 +1,5 @@
1# (beta)torch.distributed.is_hccl_available1# (beta)torch.distributed.is_hccl_available
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -13,11 +14,13 @@
13判断HCCL通信后端是否可用,与torch.distributed.is_nccl_available类似,具体请参考[https://pytorch.org/docs/stable/distributed.html\#torch.distributed.is_nccl_available](https://pytorch.org/docs/stable/distributed.html#torch.distributed.is_nccl_available)。14判断HCCL通信后端是否可用,与torch.distributed.is_nccl_available类似,具体请参考[https://pytorch.org/docs/stable/distributed.html\#torch.distributed.is_nccl_available](https://pytorch.org/docs/stable/distributed.html#torch.distributed.is_nccl_available)。
14 15 
15## 函数原型16## 函数原型
16-```17+ 
18+```python
17torch.distributed.is_hccl_available()19torch.distributed.is_hccl_available()
18```20```
19 21 
20## 返回值说明22## 返回值说明
23+ 
21`Bool`:True为可用,False为不可用。24`Bool`:True为可用,False为不可用。
22 25 
23## 调用示例26## 调用示例
@@ -29,4 +32,4 @@ import torch_npu
29torch.distributed.is_hccl_available()32torch.distributed.is_hccl_available()
30 33 
31True34True
32-```35+```
Mdocs/zh/custom_APIs/distributed/(beta)torch_npu-distributed-all_gather_into_tensor_uneven.md+3-6
@@ -1,4 +1,5 @@
1# (beta)torch_npu.distributed.all_gather_into_tensor_uneven1# (beta)torch_npu.distributed.all_gather_into_tensor_uneven
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -12,11 +13,10 @@
12 13 
13## 函数原型14## 函数原型
14 15 
15-```16+```python
16torch_npu.distributed.all_gather_into_tensor_uneven(output, input, output_split_sizes =None, group=None, async_op=False) -> torch.distributed.distributed_c10d.Work17torch_npu.distributed.all_gather_into_tensor_uneven(output, input, output_split_sizes =None, group=None, async_op=False) -> torch.distributed.distributed_c10d.Work
17```18```
18 19 
19- 
20## 参数说明20## 参数说明
21 21 
22- **output** (`Tensor`):输出Tensor,用于接收计算数据。22- **output** (`Tensor`):输出Tensor,用于接收计算数据。
@@ -31,14 +31,12 @@ torch_npu.distributed.all_gather_into_tensor_uneven(output, input, output_split_
31 31 
32`output`的shape为所有卡上`input`的shape拼接大小。32`output`的shape为所有卡上`input`的shape拼接大小。
33 33 
34- 
35## 约束说明34## 约束说明
36 35 
37- 此接口仅可在单机场景下使用。36- 此接口仅可在单机场景下使用。
38 37 
39- `output_split_sizes`元素之和等于`output`的0维;`output_split_sizes`元素个数等于`group`的size。38- `output_split_sizes`元素之和等于`output`的0维;`output_split_sizes`元素个数等于`group`的size。
40 39 
41- 
42## 调用示例40## 调用示例
43 41 
44创建以下文件test.py并保存。42创建以下文件test.py并保存。
@@ -67,7 +65,6 @@ torch_npu.distributed.all_gather_into_tensor_uneven(
67 65 
68执行如下命令。66执行如下命令。
69 67 
70-```68+```bash
71torchrun --nproc-per-node=2 test.py69torchrun --nproc-per-node=2 test.py
72```70```
73- 
Mdocs/zh/custom_APIs/distributed/(beta)torch_npu-distributed-reinit_process_group.md+2-3
@@ -1,4 +1,5 @@
1# (beta)torch_npu.distributed.reinit_process_group1# (beta)torch_npu.distributed.reinit_process_group
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -12,7 +13,7 @@
12 13 
13## 函数原型14## 函数原型
14 15 
15-```16+```python
16torch_npu.distributed.reinit_process_group(group: Optional[ProcessGroup] = None, rebuild_link: bool = True) -> None17torch_npu.distributed.reinit_process_group(group: Optional[ProcessGroup] = None, rebuild_link: bool = True) -> None
17```18```
18 19 
@@ -29,7 +30,6 @@ torch_npu.distributed.reinit_process_group(group: Optional[ProcessGroup] = None,
29 30 
30输入要确保是一个有效的device。31输入要确保是一个有效的device。
31 32 
32- 
33## 调用示例33## 调用示例
34 34 
35```python35```python
@@ -61,4 +61,3 @@ def _multiprocess(world_size,f):
61if __name__ == '__main__':61if __name__ == '__main__':
62 _multiprocess(4, _do_allreduce)62 _multiprocess(4, _do_allreduce)
63```63```
64- 
Mdocs/zh/custom_APIs/menu_Pytorch_API.md+1-1
@@ -1,4 +1,5 @@
1# 自定义API参考1# 自定义API参考
2+ 
2- [概述](./overview.md)3- [概述](./overview.md)
3- [Python接口](./Python_interface.md)4- [Python接口](./Python_interface.md)
4 - [torch_npu](./torch_npu/torch_npu.md)5 - [torch_npu](./torch_npu/torch_npu.md)
@@ -390,4 +391,3 @@
390- [废弃API列表](./scrap_API.md)391- [废弃API列表](./scrap_API.md)
391- [附录](./appendix.md)392- [附录](./appendix.md)
392 - [添加二进制黑名单示例](./blacklist.md)393 - [添加二进制黑名单示例](./blacklist.md)
393- 
Mdocs/zh/custom_APIs/scrap_API.md+36-41
@@ -9,190 +9,185 @@
9</th>9</th>
10</tr>10</tr>
11</thead>11</thead>
12-<tbody><tr id="row17311114115553"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p153112414559"><a name="p153112414559"></a><a name="p153112414559"></a><a href="(beta)torch_npu-copy_memory_.md">torch_npu.copy_memory_</a></p>12+<tbody><tr id="row17311114115553"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p153112414559"><a name="p153112414559"></a><a name="p153112414559"></a><a href="./torch_npu/(beta)torch_npu-copy_memory_.md">torch_npu.copy_memory_</a></p>
13</td>13</td>
14<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p186351481316"><a name="p186351481316"></a><a name="p186351481316"></a>该接口计划废弃,可以使用torch.Tensor.copy_接口进行替换。</p>14<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p186351481316"><a name="p186351481316"></a><a name="p186351481316"></a>该接口计划废弃,可以使用torch.Tensor.copy_接口进行替换。</p>
15</td>15</td>
16</tr>16</tr>
17-<tr id="row19311164145515"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p8311124175514"><a name="p8311124175514"></a><a name="p8311124175514"></a><a href="(beta)torch_npu-empty_with_format.md">torch_npu.empty_with_format</a></p>17+<tr id="row19311164145515"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p8311124175514"><a name="p8311124175514"></a><a name="p8311124175514"></a><a href="./torch_npu/(beta)torch_npu-empty_with_format.md">torch_npu.empty_with_format</a></p>
18</td>18</td>
19<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p129949526516"><a name="p129949526516"></a><a name="p129949526516"></a>该接口计划废弃,可以使用torch.empty接口进行替换。</p>19<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p129949526516"><a name="p129949526516"></a><a name="p129949526516"></a>该接口计划废弃,可以使用torch.empty接口进行替换。</p>
20</td>20</td>
21</tr>21</tr>
22-<tr id="row18312341155517"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p173125413551"><a name="p173125413551"></a><a name="p173125413551"></a><a href="(beta)torch_npu-npu_apply_adam.md">torch_npu.npu_apply_adam</a></p>22+<tr id="row18312341155517"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p173125413551"><a name="p173125413551"></a><a name="p173125413551"></a><a href="./torch_npu/(beta)torch_npu-npu_apply_adam.md">torch_npu.npu_apply_adam</a></p>
23</td>23</td>
24<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1899612582515"><a name="p1899612582515"></a><a name="p1899612582515"></a>该接口计划废弃,可以使用torch.optim.Adam接口进行替换。</p>24<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1899612582515"><a name="p1899612582515"></a><a name="p1899612582515"></a>该接口计划废弃,可以使用torch.optim.Adam接口进行替换。</p>
25</td>25</td>
26</tr>26</tr>
27-<tr id="row7674950125010"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p367510503501"><a name="p367510503501"></a><a name="p367510503501"></a><a href="(beta)torch_npu-npu_broadcast.md">torch_npu.npu_broadcast</a></p>27+<tr id="row7674950125010"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p367510503501"><a name="p367510503501"></a><a name="p367510503501"></a><a href="./torch_npu/(beta)torch_npu-npu_broadcast.md">torch_npu.npu_broadcast</a></p>
28</td>28</td>
29<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p121121665216"><a name="p121121665216"></a><a name="p121121665216"></a>该接口计划废弃,可以使用torch.broadcast_to接口进行替换。</p>29<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p121121665216"><a name="p121121665216"></a><a name="p121121665216"></a>该接口计划废弃,可以使用torch.broadcast_to接口进行替换。</p>
30</td>30</td>
31</tr>31</tr>
32-<tr id="row183121141175518"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p2031274114556"><a name="p2031274114556"></a><a name="p2031274114556"></a><a href="(beta)torch_npu-npu_conv_transpose2d.md">torch_npu.npu_conv_transpose2d</a></p>32+<tr id="row183121141175518"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p2031274114556"><a name="p2031274114556"></a><a name="p2031274114556"></a><a href="./torch_npu/(beta)torch_npu-npu_conv_transpose2d.md">torch_npu.npu_conv_transpose2d</a></p>
33</td>33</td>
34<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p8312194165516"><a name="p8312194165516"></a><a name="p8312194165516"></a>该接口计划废弃,可以使用torch.nn.functional.conv_transpose2d接口进行替换。</p>34<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p8312194165516"><a name="p8312194165516"></a><a name="p8312194165516"></a>该接口计划废弃,可以使用torch.nn.functional.conv_transpose2d接口进行替换。</p>
35</td>35</td>
36</tr>36</tr>
37-<tr id="row123123412554"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p9312104111558"><a name="p9312104111558"></a><a name="p9312104111558"></a><a href="(beta)torch_npu-npu_conv2d.md">torch_npu.npu_conv2d</a></p>37+<tr id="row123123412554"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p9312104111558"><a name="p9312104111558"></a><a name="p9312104111558"></a><a href="./torch_npu/(beta)torch_npu-npu_conv2d.md">torch_npu.npu_conv2d</a></p>
38</td>38</td>
39<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1242918202521"><a name="p1242918202521"></a><a name="p1242918202521"></a>该接口计划废弃,可以使用torch.nn.functional.conv2d接口进行替换。</p>39<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1242918202521"><a name="p1242918202521"></a><a name="p1242918202521"></a>该接口计划废弃,可以使用torch.nn.functional.conv2d接口进行替换。</p>
40</td>40</td>
41</tr>41</tr>
42-<tr id="row12312741155519"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p83121841105519"><a name="p83121841105519"></a><a name="p83121841105519"></a><a href="(beta)torch_npu-npu_convolution.md">torch_npu.npu_convolution</a></p>42+<tr id="row12312741155519"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p83121841105519"><a name="p83121841105519"></a><a name="p83121841105519"></a><a href="./torch_npu/(beta)torch_npu-npu_convolution.md">torch_npu.npu_convolution</a></p>
43</td>43</td>
44<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p9312041115515"><a name="p9312041115515"></a><a name="p9312041115515"></a>该接口计划废弃,可以使用torch.nn.functional.conv2d或torch.nn.functional.conv3d接口进行替换。</p>44<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p9312041115515"><a name="p9312041115515"></a><a name="p9312041115515"></a>该接口计划废弃,可以使用torch.nn.functional.conv2d或torch.nn.functional.conv3d接口进行替换。</p>
45</td>45</td>
46</tr>46</tr>
47-<tr id="row78665276418"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p48674279413"><a name="p48674279413"></a><a name="p48674279413"></a><a href="(beta)torch_npu-npu_convolution_transpose.md">torch_npu.npu_convolution_transpose</a></p>47+<tr id="row78665276418"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p48674279413"><a name="p48674279413"></a><a name="p48674279413"></a><a href="./torch_npu/(beta)torch_npu-npu_convolution_transpose.md">torch_npu.npu_convolution_transpose</a></p>
48</td>48</td>
49<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p146301141145217"><a name="p146301141145217"></a><a name="p146301141145217"></a>该接口计划废弃,可以使用torch.nn.functional.conv_transpose2d或torch.nn.functional.conv_transpose3d接口进行替换。</p>49<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p146301141145217"><a name="p146301141145217"></a><a name="p146301141145217"></a>该接口计划废弃,可以使用torch.nn.functional.conv_transpose2d或torch.nn.functional.conv_transpose3d接口进行替换。</p>
50</td>50</td>
51</tr>51</tr>
52-<tr id="row48671427164110"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p8867627114111"><a name="p8867627114111"></a><a name="p8867627114111"></a><a href="(beta)torch_npu-npu_dtype_cast.md">torch_npu.npu_dtype_cast</a></p>52+<tr id="row48671427164110"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p8867627114111"><a name="p8867627114111"></a><a name="p8867627114111"></a><a href="./torch_npu/(beta)torch_npu-npu_dtype_cast.md">torch_npu.npu_dtype_cast</a></p>
53</td>53</td>
54<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p15123104845216"><a name="p15123104845216"></a><a name="p15123104845216"></a>该接口计划废弃,可以使用torch.to接口进行替换。</p>54<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p15123104845216"><a name="p15123104845216"></a><a name="p15123104845216"></a>该接口计划废弃,可以使用torch.to接口进行替换。</p>
55</td>55</td>
56</tr>56</tr>
57-<tr id="row186752718417"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p4868162714410"><a name="p4868162714410"></a><a name="p4868162714410"></a><a href="(beta)torch_npu-npu_gru.md">torch_npu.npu_gru</a></p>57+<tr id="row186752718417"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p4868162714410"><a name="p4868162714410"></a><a name="p4868162714410"></a><a href="./torch_npu/(beta)torch_npu-npu_gru.md">torch_npu.npu_gru</a></p>
58</td>58</td>
59<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p45973015310"><a name="p45973015310"></a><a name="p45973015310"></a>该接口计划废弃,可以使用torch.gru接口进行替换。</p>59<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p45973015310"><a name="p45973015310"></a><a name="p45973015310"></a>该接口计划废弃,可以使用torch.gru接口进行替换。</p>
60</td>60</td>
61</tr>61</tr>
62-<tr id="row6868202716414"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p16868727194118"><a name="p16868727194118"></a><a name="p16868727194118"></a><a href="(beta)torch_npu-npu_layer_norm_eval.md">torch_npu.npu_layer_norm_eval</a></p>62+<tr id="row6868202716414"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p16868727194118"><a name="p16868727194118"></a><a name="p16868727194118"></a><a href="./torch_npu/(beta)torch_npu-npu_layer_norm_eval.md">torch_npu.npu_layer_norm_eval</a></p>
63</td>63</td>
64<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p186842764114"><a name="p186842764114"></a><a name="p186842764114"></a>该接口计划废弃,可以使用torch.nn.functional.layer_norm接口进行替换。</p>64<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p186842764114"><a name="p186842764114"></a><a name="p186842764114"></a>该接口计划废弃,可以使用torch.nn.functional.layer_norm接口进行替换。</p>
65</td>65</td>
66</tr>66</tr>
67-<tr id="row786882717416"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p1586882734118"><a name="p1586882734118"></a><a name="p1586882734118"></a><a href="(beta)torch_npu-npu_min.md">torch_npu.npu_min</a></p>67+<tr id="row786882717416"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p1586882734118"><a name="p1586882734118"></a><a name="p1586882734118"></a><a href="./torch_npu/(beta)torch_npu-npu_min.md">torch_npu.npu_min</a></p>
68</td>68</td>
69<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p547813112537"><a name="p547813112537"></a><a name="p547813112537"></a>该接口计划废弃,可以使用torch.min接口进行替换。</p>69<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p547813112537"><a name="p547813112537"></a><a name="p547813112537"></a>该接口计划废弃,可以使用torch.min接口进行替换。</p>
70</td>70</td>
71</tr>71</tr>
72-<tr id="row9330173494118"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p933013416416"><a name="p933013416416"></a><a name="p933013416416"></a><a href="(beta)torch_npu-npu_mish.md">torch_npu.npu_mish</a></p>72+<tr id="row9330173494118"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p933013416416"><a name="p933013416416"></a><a name="p933013416416"></a><a href="./torch_npu/(beta)torch_npu-npu_mish.md">torch_npu.npu_mish</a></p>
73</td>73</td>
74<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p202681337195317"><a name="p202681337195317"></a><a name="p202681337195317"></a>该接口计划废弃,可以使用torch.nn.functional.mish接口进行替换。</p>74<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p202681337195317"><a name="p202681337195317"></a><a name="p202681337195317"></a>该接口计划废弃,可以使用torch.nn.functional.mish接口进行替换。</p>
75</td>75</td>
76</tr>76</tr>
77-<tr id="row0331103454113"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p19331734194116"><a name="p19331734194116"></a><a name="p19331734194116"></a><a href="(beta)torch_npu-npu_nms_rotated.md">torch_npu.npu_nms_rotated</a></p>77+<tr id="row0331103454113"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p19331734194116"><a name="p19331734194116"></a><a name="p19331734194116"></a><a href="./torch_npu/(beta)torch_npu-npu_nms_rotated.md">torch_npu.npu_nms_rotated</a></p>
78</td>78</td>
79<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p208131543205314"><a name="p208131543205314"></a><a name="p208131543205314"></a>该接口计划废弃,可以参考<a href="https://gitcode.com/Ascend/op-plugin/blob/7.3.0/test/test_base_ops/test_nms_rotated.py" target="_blank" rel="noopener noreferrer">小算子拼接方案</a>进行替换。</p>79<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p208131543205314"><a name="p208131543205314"></a><a name="p208131543205314"></a>该接口计划废弃,可以参考<a href="https://gitcode.com/Ascend/op-plugin/blob/7.3.0/test/test_base_ops/test_nms_rotated.py" target="_blank" rel="noopener noreferrer">小算子拼接方案</a>进行替换。</p>
80</td>80</td>
81</tr>81</tr>
82-<tr id="row43319341418"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p63313347411"><a name="p63313347411"></a><a name="p63313347411"></a><a href="(beta)torch_npu-npu_ptiou.md">torch_npu.npu_ptiou</a></p>82+<tr id="row43319341418"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p63313347411"><a name="p63313347411"></a><a name="p63313347411"></a><a href="./torch_npu/(beta)torch_npu-npu_ptiou.md">torch_npu.npu_ptiou</a></p>
83</td>83</td>
84<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p143099280139"><a name="p143099280139"></a><a name="p143099280139"></a>该接口计划废弃,可以使用torch_npu.npu_iou接口进行替换。</p>84<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p143099280139"><a name="p143099280139"></a><a name="p143099280139"></a>该接口计划废弃,可以使用torch_npu.npu_iou接口进行替换。</p>
85</td>85</td>
86</tr>86</tr>
87-<tr id="row333133464112"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p1533111343412"><a name="p1533111343412"></a><a name="p1533111343412"></a><a href="(beta)torch_npu-npu_reshape.md">torch_npu.npu_reshape</a></p>87+<tr id="row333133464112"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p1533111343412"><a name="p1533111343412"></a><a name="p1533111343412"></a><a href="./torch_npu/(beta)torch_npu-npu_reshape.md">torch_npu.npu_reshape</a></p>
88</td>88</td>
89<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p11331034154114"><a name="p11331034154114"></a><a name="p11331034154114"></a>该接口计划废弃,可以使用torch.reshape接口进行替换。</p>89<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p11331034154114"><a name="p11331034154114"></a><a name="p11331034154114"></a>该接口计划废弃,可以使用torch.reshape接口进行替换。</p>
90</td>90</td>
91</tr>91</tr>
92-<tr id="row1433223417411"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p183324349416"><a name="p183324349416"></a><a name="p183324349416"></a><a href="(beta)torch_npu-npu_silu.md">torch_npu.npu_silu</a></p>92+<tr id="row1433223417411"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p183324349416"><a name="p183324349416"></a><a name="p183324349416"></a><a href="./torch_npu/(beta)torch_npu-npu_silu.md">torch_npu.npu_silu</a></p>
93</td>93</td>
94<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p11526133155418"><a name="p11526133155418"></a><a name="p11526133155418"></a>该接口计划废弃,可以使用torch.nn.functional.silu接口进行替换。</p>94<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p11526133155418"><a name="p11526133155418"></a><a name="p11526133155418"></a>该接口计划废弃,可以使用torch.nn.functional.silu接口进行替换。</p>
95</td>95</td>
96</tr>96</tr>
97-<tr id="row13332734114113"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p183327342419"><a name="p183327342419"></a><a name="p183327342419"></a><a href="(beta)torch_npu-npu_sort_v2.md">torch_npu.npu_sort_v2</a></p>97+<tr id="row13332734114113"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p183327342419"><a name="p183327342419"></a><a name="p183327342419"></a><a href="./torch_npu/(beta)torch_npu-npu_sort_v2.md">torch_npu.npu_sort_v2</a></p>
98</td>98</td>
99<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p128654045419"><a name="p128654045419"></a><a name="p128654045419"></a>该接口计划废弃,可以使用torch.sort接口进行替换。</p>99<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p128654045419"><a name="p128654045419"></a><a name="p128654045419"></a>该接口计划废弃,可以使用torch.sort接口进行替换。</p>
100</td>100</td>
101</tr>101</tr>
102-<tr id="row533243424115"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p183321434184119"><a name="p183321434184119"></a><a name="p183321434184119"></a><a href="(beta)torch_npu-one_.md">torch_npu.one_</a></p>102+<tr id="row533243424115"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p183321434184119"><a name="p183321434184119"></a><a name="p183321434184119"></a><a href="./torch_npu/(beta)torch_npu-one_.md">torch_npu.one_</a></p>
103</td>103</td>
104<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1533283474110"><a name="p1533283474110"></a><a name="p1533283474110"></a>该接口计划废弃,可以使用torch.fill_或torch.ones_like接口进行替换。</p>104<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1533283474110"><a name="p1533283474110"></a><a name="p1533283474110"></a>该接口计划废弃,可以使用torch.fill_或torch.ones_like接口进行替换。</p>
105</td>105</td>
106</tr>106</tr>
107-<tr id="row14332534124118"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p9332173420417"><a name="p9332173420417"></a><a name="p9332173420417"></a><a href="(beta)torch_npu-contrib-DCNv2.md">torch_npu.contrib.DCNv2</a></p>107+<tr id="row14332534124118"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p9332173420417"><a name="p9332173420417"></a><a name="p9332173420417"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-DCNv2.md">torch_npu.contrib.DCNv2</a></p>
108</td>108</td>
109<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p16651085550"><a name="p16651085550"></a><a name="p16651085550"></a>该接口计划废弃,可以使用torch_npu.contrib.ModulationDeformCon接口进行替换。</p>109<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p16651085550"><a name="p16651085550"></a><a name="p16651085550"></a>该接口计划废弃,可以使用torch_npu.contrib.ModulationDeformCon接口进行替换。</p>
110</td>110</td>
111</tr>111</tr>
112-<tr id="row2333123474117"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p733363404114"><a name="p733363404114"></a><a name="p733363404114"></a><a href="(beta)torch_npu-contrib-BiLSTM.md">torch_npu.contrib.BiLSTM</a></p>112+<tr id="row2333123474117"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p733363404114"><a name="p733363404114"></a><a name="p733363404114"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-BiLSTM.md">torch_npu.contrib.BiLSTM</a></p>
113</td>113</td>
114<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p191653147552"><a name="p191653147552"></a><a name="p191653147552"></a>该接口计划废弃,可以参考<a href="https://gitee.com/ascend/ModelZoo-PyTorch/blob/732cb7fc5ab59249ae62a905c0d43400a8250da7/PyTorch/contrib/audio/deepspeech/deepspeech_pytorch/bidirectional_lstm.py#L18" target="_blank" rel="noopener noreferrer">小算子拼接方案</a>进行替换。</p>114<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p191653147552"><a name="p191653147552"></a><a name="p191653147552"></a>该接口计划废弃,可以参考<a href="https://gitee.com/ascend/ModelZoo-PyTorch/blob/732cb7fc5ab59249ae62a905c0d43400a8250da7/PyTorch/contrib/audio/deepspeech/deepspeech_pytorch/bidirectional_lstm.py#L18" target="_blank" rel="noopener noreferrer">小算子拼接方案</a>进行替换。</p>
115</td>115</td>
116</tr>116</tr>
117-<tr id="row1846663914412"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p15466439164114"><a name="p15466439164114"></a><a name="p15466439164114"></a><a href="(beta)torch_npu-contrib-Swish.md">torch_npu.contrib.Swish</a></p>117+<tr id="row1846663914412"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p15466439164114"><a name="p15466439164114"></a><a name="p15466439164114"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-Swish.md">torch_npu.contrib.Swish</a></p>
118</td>118</td>
119<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p74504207559"><a name="p74504207559"></a><a name="p74504207559"></a>该接口计划废弃,可以使用torch_npu.contrib.ModulationDeformCon接口进行替换。</p>119<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p74504207559"><a name="p74504207559"></a><a name="p74504207559"></a>该接口计划废弃,可以使用torch_npu.contrib.ModulationDeformCon接口进行替换。</p>
120</td>120</td>
121</tr>121</tr>
122-<tr id="row13466103954110"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p5467143917412"><a name="p5467143917412"></a><a name="p5467143917412"></a><a href="(beta)torch_npu-contrib-npu_giou.md">torch_npu.contrib.npu_giou</a></p>122+<tr id="row13466103954110"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p5467143917412"><a name="p5467143917412"></a><a name="p5467143917412"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-npu_giou.md">torch_npu.contrib.npu_giou</a></p>
123</td>123</td>
124<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p3480527105515"><a name="p3480527105515"></a><a name="p3480527105515"></a>该接口计划废弃,可以使用torch_npu.npu_giou接口进行替换。</p>124<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p3480527105515"><a name="p3480527105515"></a><a name="p3480527105515"></a>该接口计划废弃,可以使用torch_npu.npu_giou接口进行替换。</p>
125</td>125</td>
126</tr>126</tr>
127-<tr id="row746723920415"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p104671339184114"><a name="p104671339184114"></a><a name="p104671339184114"></a><a href="(beta)torch_npu-contrib-npu_ptiou.md">torch_npu.contrib.npu_ptiou</a></p>127+<tr id="row746723920415"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p104671339184114"><a name="p104671339184114"></a><a name="p104671339184114"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-npu_ptiou.md">torch_npu.contrib.npu_ptiou</a></p>
128</td>128</td>
129<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p119998332552"><a name="p119998332552"></a><a name="p119998332552"></a>该接口计划废弃,可以使用torch_npu.npu_iou接口进行替换。</p>129<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p119998332552"><a name="p119998332552"></a><a name="p119998332552"></a>该接口计划废弃,可以使用torch_npu.npu_iou接口进行替换。</p>
130</td>130</td>
131</tr>131</tr>
132-<tr id="row9467103913412"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p1946712394415"><a name="p1946712394415"></a><a name="p1946712394415"></a><a href="(beta)torch_npu-contrib-npu_iou.md">torch_npu.contrib.npu_iou</a></p>132+<tr id="row9467103913412"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p1946712394415"><a name="p1946712394415"></a><a name="p1946712394415"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-npu_iou.md">torch_npu.contrib.npu_iou</a></p>
133</td>133</td>
134<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1261323935510"><a name="p1261323935510"></a><a name="p1261323935510"></a>该接口计划废弃,可以使用torch_npu.npu_iou接口进行替换。</p>134<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1261323935510"><a name="p1261323935510"></a><a name="p1261323935510"></a>该接口计划废弃,可以使用torch_npu.npu_iou接口进行替换。</p>
135</td>135</td>
136</tr>136</tr>
137-<tr id="row174671139164110"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p84677395415"><a name="p84677395415"></a><a name="p84677395415"></a><a href="(beta)torch_npu-contrib-function-npu_diou.md">torch_npu.contrib.function.npu_diou</a></p>137+<tr id="row174671139164110"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p84677395415"><a name="p84677395415"></a><a name="p84677395415"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-function-npu_diou.md">torch_npu.contrib.function.npu_diou</a></p>
138</td>138</td>
139<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p3545164812557"><a name="p3545164812557"></a><a name="p3545164812557"></a>该接口计划废弃,可以使用torch_npu.npu_diou接口进行替换。</p>139<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p3545164812557"><a name="p3545164812557"></a><a name="p3545164812557"></a>该接口计划废弃,可以使用torch_npu.npu_diou接口进行替换。</p>
140</td>140</td>
141</tr>141</tr>
142-<tr id="row13467173915415"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p24671739124114"><a name="p24671739124114"></a><a name="p24671739124114"></a><a href="(beta)torch_npu-contrib-function-npu_ciou.md">torch_npu.contrib.function.npu_ciou</a></p>142+<tr id="row13467173915415"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p24671739124114"><a name="p24671739124114"></a><a name="p24671739124114"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-function-npu_ciou.md">torch_npu.contrib.function.npu_ciou</a></p>
143</td>143</td>
144<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1561015645519"><a name="p1561015645519"></a><a name="p1561015645519"></a>该接口计划废弃,可以使用torch_npu.npu_ciou接口进行替换。</p>144<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p1561015645519"><a name="p1561015645519"></a><a name="p1561015645519"></a>该接口计划废弃,可以使用torch_npu.npu_ciou接口进行替换。</p>
145</td>145</td>
146</tr>146</tr>
147-<tr id="row046716399413"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p164671839164116"><a name="p164671839164116"></a><a name="p164671839164116"></a><a href="(beta)torch_npu-contrib-module-Mish.md">torch_npu.contrib.module.Mish</a></p>147+<tr id="row046716399413"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p164671839164116"><a name="p164671839164116"></a><a name="p164671839164116"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-module-Mish.md">torch_npu.contrib.module.Mish</a></p>
148</td>148</td>
149<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p763677195619"><a name="p763677195619"></a><a name="p763677195619"></a>该接口计划废弃,可以使用torch.nn.Mish接口进行替换。</p>149<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p763677195619"><a name="p763677195619"></a><a name="p763677195619"></a>该接口计划废弃,可以使用torch.nn.Mish接口进行替换。</p>
150</td>150</td>
151</tr>151</tr>
152-<tr id="row124676390414"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p7467539124114"><a name="p7467539124114"></a><a name="p7467539124114"></a><a href="(beta)torch_npu-contrib-module-SiLU.md">torch_npu.contrib.module.SiLU</a></p>152+<tr id="row124676390414"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p7467539124114"><a name="p7467539124114"></a><a name="p7467539124114"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-module-SiLU.md">torch_npu.contrib.module.SiLU</a></p>
153</td>153</td>
154<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p10467103919418"><a name="p10467103919418"></a><a name="p10467103919418"></a>该接口计划废弃,可以使用torch.nn.SiLU接口进行替换。</p>154<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p10467103919418"><a name="p10467103919418"></a><a name="p10467103919418"></a>该接口计划废弃,可以使用torch.nn.SiLU接口进行替换。</p>
155</td>155</td>
156</tr>156</tr>
157-<tr id="row11468103954112"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p154681639164115"><a name="p154681639164115"></a><a name="p154681639164115"></a><a href="(beta)torch_npu-contrib-module-FusedColorJitter.md">torch_npu.contrib.module.FusedColorJitter</a></p>157+<tr id="row11468103954112"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p154681639164115"><a name="p154681639164115"></a><a name="p154681639164115"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-module-FusedColorJitter.md">torch_npu.contrib.module.FusedColorJitter</a></p>
158</td>158</td>
159<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p918621912569"><a name="p918621912569"></a><a name="p918621912569"></a>该接口计划废弃,可以使用torchvision.transforms.ColorJitter接口进行替换。</p>159<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p918621912569"><a name="p918621912569"></a><a name="p918621912569"></a>该接口计划废弃,可以使用torchvision.transforms.ColorJitter接口进行替换。</p>
160</td>160</td>
161</tr>161</tr>
162-<tr id="row44381047182110"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p147111718151717"><a name="p147111718151717"></a><a name="p147111718151717"></a><a href="torch_npu-contrib-module-LinearA8W8Quant.md">torch_npu.contrib.module.LinearA8W8Quant</a></p>162+<tr id="row44381047182110"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p147111718151717"><a name="p147111718151717"></a><a name="p147111718151717"></a><a href="./torch_npu-contrib/torch_npu-contrib-module-LinearA8W8Quant.md">torch_npu.contrib.module.LinearA8W8Quant</a></p>
163</td>163</td>
164<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p12865714111912"><a name="p12865714111912"></a><a name="p12865714111912"></a>该接口计划废弃,可以使用torch_npu.contrib.module.LinearQuant接口进行替换。</p>164<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p12865714111912"><a name="p12865714111912"></a><a name="p12865714111912"></a>该接口计划废弃,可以使用torch_npu.contrib.module.LinearQuant接口进行替换。</p>
165</td>165</td>
166</tr>166</tr>
167-<tr id="row1597725217179"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="(beta)torch_npu-contrib-npu_fused_attention_with_layernorm.md">torch_npu.contrib.npu_fused_attention_with_layernorm</a></p>167+<tr id="row1597725217179"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="./torch_npu-contrib/(beta)torch_npu-contrib-npu_fused_attention_with_layernorm.md">torch_npu.contrib.npu_fused_attention_with_layernorm</a></p>
168</td>168</td>
169<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p169771352131719"><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,可以使用torch_npu.npu_fusion_attention与torch.nn.LayerNorm接口进行替换。</p>169<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p169771352131719"><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,可以使用torch_npu.npu_fusion_attention与torch.nn.LayerNorm接口进行替换。</p>
170</td>170</td>
171</tr>171</tr>
172-</tr>172+<tr><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="./torch_npu/(beta)torch_npu-npu_conv3d.md">torch_npu.npu_conv3d</a></p>
173-<tr><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="(beta)torch_npu-npu_conv3d.md">torch_npu.npu_conv3d</a></p>
174</td>173</td>
175<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,可以使用`torch.nn.functional.conv3d`接口进行替换。</p>174<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,可以使用`torch.nn.functional.conv3d`接口进行替换。</p>
176</td>175</td>
177</tr>176</tr>
178-</tr>177+<tr><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="./torch_npu/(beta)torch_npu-npu_bmmV2.md">torch_npu.npu_bmmV2</a></p>
179-<tr><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="(beta)torch_npu-npu_bmmV2.md">torch_npu.npu_bmmV2</a></p>
180</td>178</td>
181<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,可以使用`torch.bmm``torch.view`接口进行替换。</p>179<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,可以使用`torch.bmm``torch.view`接口进行替换。</p>
182</td>180</td>
183</tr>181</tr>
184-</tr>182+<tr><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="(./torch_npu/(beta)torch_npu-npu_confusion_transpose.md">torch_npu.npu_confusion_transpose</a></p>
185-<tr><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="(beta)torch_npu-npu_confusion_transpose.md">torch_npu.npu_confusion_transpose</a></p>
186</td>183</td>
187<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,可以使用`torch.view``torch.permute`接口进行替换。</p>184<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,可以使用`torch.view``torch.permute`接口进行替换。</p>
188</td>185</td>
189</tr>186</tr>
190-</tr>187+<tr id="row1597725217179"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p11977252111717"><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="./torch_npu-npu/torch_npu-npu-ExternalEvent().reset().md">torch_npu-npu-ExternalEvent().reset()</a></p>
191-<tr id="row1597725217179"><td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.1 "><p id="p11977252111717"><a name="p11977252111717"></a><a name="p11977252111717"></a><a href="torch_npu-npu-ExternalEvent().reset().md">torch_npu-npu-ExternalEvent().reset()</a></p>
192</td>188</td>
193<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p169771352131719"><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,torch_npu.npu.ExternalEvent().wait()会自动复位Event,不推荐调用本接口手动复位Event。</p>189<td class="cellrowborder" valign="top" width="50%" headers="mcps1.2.3.1.2 "><p id="p169771352131719"><a name="p169771352131719"></a><a name="p169771352131719"></a>该接口计划废弃,torch_npu.npu.ExternalEvent().wait()会自动复位Event,不推荐调用本接口手动复位Event。</p>
194</td>190</td>
195</tr>191</tr>
196</tbody>192</tbody>
197</table>193</table>
198- 
Mdocs/zh/custom_APIs/torch_npu-contrib/torch_npu-contrib-module-LinearA8W8Quant.md+2-4
@@ -11,14 +11,13 @@
11| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |11| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |
12| <term>Atlas 推理系列产品</term> | √ |12| <term>Atlas 推理系列产品</term> | √ |
13 13 
14- 
15## 功能说明14## 功能说明
16 15 
17LinearA8W8Quant是对torch_npu.npu_quant_matmul接口的封装类,完成A8W8量化算子的矩阵乘计算。16LinearA8W8Quant是对torch_npu.npu_quant_matmul接口的封装类,完成A8W8量化算子的矩阵乘计算。
18 17 
19## 函数原型18## 函数原型
20 19 
21-```20+```python
22torch_npu.contrib.module.LinearA8W8Quant(in_features, out_features, *, bias=True, offset=False, pertoken_scale=False, output_dtype=None)21torch_npu.contrib.module.LinearA8W8Quant(in_features, out_features, *, bias=True, offset=False, pertoken_scale=False, output_dtype=None)
23```22```
24 23 
@@ -60,6 +59,7 @@ torch_npu.contrib.module.LinearA8W8Quant(in_features, out_features, *, bias=True
60 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持输入`int8``float16``bfloat16`59 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持输入`int8``float16``bfloat16`
61 60 
62## 返回值说明61## 返回值说明
62+ 
63`Tensor`63`Tensor`
64 64 
65代表量化matmul的计算结果:65代表量化matmul的计算结果:
@@ -216,7 +216,6 @@ torch_npu.contrib.module.LinearA8W8Quant(in_features, out_features, *, bias=True
216 </tbody>216 </tbody>
217 </table>217 </table>
218 218 
219- 
220## 调用示例219## 调用示例
221 220 
222- 单算子模式调用221- 单算子模式调用
@@ -282,4 +281,3 @@ torch_npu.contrib.module.LinearA8W8Quant(in_features, out_features, *, bias=True
282 model = torch.compile(model, backend=npu_backend, dynamic=False)281 model = torch.compile(model, backend=npu_backend, dynamic=False)
283 output = model(x1)282 output = model(x1)
284 ```283 ```
285- 
Mdocs/zh/custom_APIs/torch_npu-contrib/torch_npu-contrib-module-LinearQuant.md+1-2
@@ -14,7 +14,7 @@ LinearQuant是对torch_npu.npu_quant_matmul接口的封装类,完成A8W8、A4W
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.contrib.module.LinearQuant(in_features, out_features, *, bias=True, offset=False, pertoken_scale=False, device=None, dtype=None, output_dtype=None)18torch_npu.contrib.module.LinearQuant(in_features, out_features, *, bias=True, offset=False, pertoken_scale=False, device=None, dtype=None, output_dtype=None)
19```19```
20 20 
@@ -349,4 +349,3 @@ torch_npu.contrib.module.LinearQuant(in_features, out_features, *, bias=True, of
349 model = torch.compile(model, backend=npu_backend, dynamic=False)349 model = torch.compile(model, backend=npu_backend, dynamic=False)
350 output = model(x1)350 output = model(x1)
351 ```351 ```
352- 
Mdocs/zh/custom_APIs/torch_npu-contrib/torch_npu-contrib-module-LinearWeightQuant.md+4-5
@@ -16,7 +16,7 @@ LinearWeightQuant是对torch_npu.npu_weight_quant_batchmatmul接口的封装类
16 16 
17## 函数原型17## 函数原型
18 18 
19-```19+```python
20torch_npu.contrib.module.LinearWeightQuant(in_features, out_features, bias=True, device=None, dtype=None, antiquant_offset=False, quant_scale=False, quant_offset=False, antiquant_group_size=0, inner_precise=0)20torch_npu.contrib.module.LinearWeightQuant(in_features, out_features, bias=True, device=None, dtype=None, antiquant_offset=False, quant_scale=False, quant_offset=False, antiquant_group_size=0, inner_precise=0)
21```21```
22 22 
@@ -45,7 +45,7 @@ torch_npu.contrib.module.LinearWeightQuant(in_features, out_features, bias=True,
45## 变量说明45## 变量说明
46 46 
47- **weight**`Tensor`):即矩阵乘中的weight。数据格式支持$ND$、FRACTAL_NZ,支持非连续的Tensor,支持输入维度为两维(N, K)。47- **weight**`Tensor`):即矩阵乘中的weight。数据格式支持$ND$、FRACTAL_NZ,支持非连续的Tensor,支持输入维度为两维(N, K)。
48- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品</term>:数据类型支持`int8`、`int32`(通过`int32`承载`int4`的输入,可以参考[torch_npu.npu_convert_weight_to_int4pack](torch_npu-npu_convert_weight_to_int4pack.md)的调用示例)。48+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品</term>:数据类型支持`int8`、`int32`(通过`int32`承载`int4`的输入,可以参考[torch_npu.npu_convert_weight_to_int4pack](../torch_npu/torch_npu-npu_convert_weight_to_int4pack.md)的调用示例)。
49 - <term>Atlas 推理系列产品</term>:数据类型支持`int8`。weight FRACTAL_NZ格式只在图模式有效,依赖接口torchair.experimental.inference.use_internal_format_weight完成数据格式从ND到FRACTAL_NZ转换,可参考[调用示例](#section00001)。49 - <term>Atlas 推理系列产品</term>:数据类型支持`int8`。weight FRACTAL_NZ格式只在图模式有效,依赖接口torchair.experimental.inference.use_internal_format_weight完成数据格式从ND到FRACTAL_NZ转换,可参考[调用示例](#section00001)。
50 50 
51- **antiquant_scale**(`Tensor`):反量化的scale,用于weight矩阵反量化。数据格式支持$ND$。支持非连续的Tensor,支持输入维度为两维(N, 1)或一维(N,)、(1,)。51- **antiquant_scale**(`Tensor`):反量化的scale,用于weight矩阵反量化。数据格式支持$ND$。支持非连续的Tensor,支持输入维度为两维(N, 1)或一维(N,)、(1,)。
@@ -53,7 +53,7 @@ torch_npu.contrib.module.LinearWeightQuant(in_features, out_features, bias=True,
53 - 若数据类型为`float16``bfloat16`,其数据类型需要和`x`保持一致。53 - 若数据类型为`float16``bfloat16`,其数据类型需要和`x`保持一致。
54 - 若数据类型为`int64`,则`x`的数据类型必须为`float16`且不带transpose输入,同时`weight`的数据类型必须为`int8`、数据格式为$ND$、带transpose输入,可参考[调用示例](#section00001)。此时只支持perchannel场景,M范围为[1, 96],且K和N要求64对齐。54 - 若数据类型为`int64`,则`x`的数据类型必须为`float16`且不带transpose输入,同时`weight`的数据类型必须为`int8`、数据格式为$ND$、带transpose输入,可参考[调用示例](#section00001)。此时只支持perchannel场景,M范围为[1, 96],且K和N要求64对齐。
55 55 
56- - <term>Atlas 推理系列产品</term> :数据类型支持`float16`,其数据类型需要和`x`保持一致。56+ - <term>Atlas 推理系列产品</term> :数据类型支持`float16`,其数据类型需要和`x`保持一致。
57 57 
58- **antiquant_offset**(`Tensor`):反量化的offset,用于weight矩阵反量化。数据格式支持$ND$。支持非连续的Tensor,支持输入维度为两维(N, 1)或一维(N,)、(1,)。58- **antiquant_offset**(`Tensor`):反量化的offset,用于weight矩阵反量化。数据格式支持$ND$。支持非连续的Tensor,支持输入维度为两维(N, 1)或一维(N,)、(1,)。
59 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品</term> :数据类型支持`float16``bfloat16``int32`。pergroup场景shape要求为(N, ceil_div(K, antiquant_group_size))。59 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品</term> :数据类型支持`float16``bfloat16``int32`。pergroup场景shape要求为(N, ceil_div(K, antiquant_group_size))。
@@ -70,6 +70,7 @@ torch_npu.contrib.module.LinearWeightQuant(in_features, out_features, bias=True,
70- **antiquant_group_size**`int`):用于控制pergroup场景下的group大小,默认为0。传入值的范围为[32, K-1]且值要求是32的倍数。<term>Atlas 推理系列产品</term> :暂不支持此参数。70- **antiquant_group_size**`int`):用于控制pergroup场景下的group大小,默认为0。传入值的范围为[32, K-1]且值要求是32的倍数。<term>Atlas 推理系列产品</term> :暂不支持此参数。
71 71 
72## 返回值说明72## 返回值说明
73+ 
73`Tensor`74`Tensor`
74 75 
75代表计算结果。当输入存在`quant_scale`时输出数据类型为`int8`,当输入不存在`quant_scale`时输出数据类型和输入`x`一致。76代表计算结果。当输入存在`quant_scale`时输出数据类型为`int8`,当输入不存在`quant_scale`时输出数据类型和输入`x`一致。
@@ -86,7 +87,6 @@ torch_npu.contrib.module.LinearWeightQuant(in_features, out_features, bias=True,
86- 如需传入`int64`数据类型的`quant_scale`,需要提前调用torch_npu.npu_trans_quant_param接口将数据类型为`float32``quant_scale``quant_offset`转换为数据类型为`int64``quant_scale`输入,可参考[调用示例](#section00001)。87- 如需传入`int64`数据类型的`quant_scale`,需要提前调用torch_npu.npu_trans_quant_param接口将数据类型为`float32``quant_scale``quant_offset`转换为数据类型为`int64``quant_scale`输入,可参考[调用示例](#section00001)。
87- 当输入`weight`为FRACTAL_NZ格式且类型为`int32`时,perchannel场景需满足`weight`为转置输入;pergroup场景需满足`x`为转置输入,`weight`为非转置输入,`antiquant_group_size`为64或128,K为`antiquant_group_size`对齐,N为64对齐。88- 当输入`weight`为FRACTAL_NZ格式且类型为`int32`时,perchannel场景需满足`weight`为转置输入;pergroup场景需满足`x`为转置输入,`weight`为非转置输入,`antiquant_group_size`为64或128,K为`antiquant_group_size`对齐,N为64对齐。
88 89 
89- 
90## 调用示例<a name="section00001"></a>90## 调用示例<a name="section00001"></a>
91 91 
92- 单算子模式调用92- 单算子模式调用
@@ -157,4 +157,3 @@ torch_npu.contrib.module.LinearWeightQuant(in_features, out_features, bias=True,
157 model = torch.compile(model, backend=npu_backend, dynamic=False)157 model = torch.compile(model, backend=npu_backend, dynamic=False)
158 out = model(x)158 out = model(x)
159 ```159 ```
160- 
Mdocs/zh/custom_APIs/torch_npu-contrib/torch_npu-contrib-module-QuantConv2d.md+2-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.module.QuantConv2d(in_channels, out_channels, kernel_size, output_dtype, stride=1, padding=0, dilation=1, groups=1, bias=True, offset=False, offset_x=0, round_mode="rint", device=None, dtype=None)22torch_npu.contrib.module.QuantConv2d(in_channels, out_channels, kernel_size, output_dtype, stride=1, padding=0, dilation=1, groups=1, bias=True, offset=False, offset_x=0, round_mode="rint", device=None, dtype=None)
23```23```
24 24 
@@ -57,6 +57,7 @@ torch_npu.contrib.module.QuantConv2d(in_channels, out_channels, kernel_size, out
57- **bias**`Tensor`):可选参数。数据类型支持`int32`,数据格式支持$ND$,shape支持1维(n,),n与`weight``out_channels`一致。57- **bias**`Tensor`):可选参数。数据类型支持`int32`,数据格式支持$ND$,shape支持1维(n,),n与`weight``out_channels`一致。
58 58 
59## 输出说明59## 输出说明
60+ 
60`Tensor`61`Tensor`
61 62 
62代表QuantConv2d的计算结果:63代表QuantConv2d的计算结果:
@@ -102,4 +103,3 @@ with torch.no_grad():
102 output = static_graph_model(fmap)103 output = static_graph_model(fmap)
103print("static graph result: ", output)104print("static graph result: ", output)
104```105```
105- 
Mdocs/zh/custom_APIs/torch_npu-contrib/torch_npu-contrib.md+1-1
@@ -1 +1 @@
1-# torch_npu.contrib1+# torch_npu.contrib
Mdocs/zh/custom_APIs/torch_npu-contrib/torch_npu-contrib_list.md+0-1
@@ -295,4 +295,3 @@
295</tr>295</tr>
296</tbody>296</tbody>
297</table>297</table>
298- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-BiLSTM.md+1-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.BiLSTM(input_size, hidden_size)22torch_npu.contrib.BiLSTM(input_size, hidden_size)
23```23```
24 24 
@@ -36,4 +36,3 @@ torch_npu.contrib.BiLSTM(input_size, hidden_size)
36>>> input_tensor = torch.randn(26, 2560, 512).npu()36>>> input_tensor = torch.randn(26, 2560, 512).npu()
37>>> output = r(input_tensor)37>>> output = r(input_tensor)
38```38```
39- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-DCNv2.md+1-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.DCNv2(in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, deformable_groups=1, bias=True, pack=True)22torch_npu.contrib.DCNv2(in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, deformable_groups=1, bias=True, pack=True)
23```23```
24 24 
@@ -50,4 +50,3 @@ ModulationDeformConv仅实现fp32数据类型下的操作。注意,conv_offset
50>>> output = model(x)50>>> output = model(x)
51>>> output.sum().backward()51>>> output.sum().backward()
52```52```
53- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-NpuFairseqDropout.md+2-2
@@ -15,7 +15,7 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.NpuFairseqDropout(p, module_name=None)19torch_npu.contrib.NpuFairseqDropout(p, module_name=None)
20```20```
21 21 
@@ -26,4 +26,4 @@ torch_npu.contrib.NpuFairseqDropout(p, module_name=None)
26 26 
27## 约束说明27## 约束说明
28 28 
29-不支持动态shape。29+不支持动态shape。
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-Prefetcher.md+1-2
@@ -11,7 +11,7 @@
11 11 
12## 函数原型12## 函数原型
13 13 
14-```14+```python
15torch_npu.contrib.Prefetcher(loader, stream=None)15torch_npu.contrib.Prefetcher(loader, stream=None)
16```16```
17 17 
@@ -23,4 +23,3 @@ NPU设备上的数据预取器,主要用于优化数据加载流程,提升
23 23 
24- **loader** (torch.utils.data.DataLoader or DataLoader like iterator):必选参数。预处理后的输入数据。24- **loader** (torch.utils.data.DataLoader or DataLoader like iterator):必选参数。预处理后的输入数据。
25- **stream** (torch.npu.Stream):可选参数,默认值为None。由于NPU内存逻辑限制,如果要在训练中重复初始化prefetcher,就需要指定一个stream来防止内存泄漏;如果prefetcher仅在训练中被初始化一次,则无需指定stream,会自动创建一个stream。25- **stream** (torch.npu.Stream):可选参数,默认值为None。由于NPU内存逻辑限制,如果要在训练中重复初始化prefetcher,就需要指定一个stream来防止内存泄漏;如果prefetcher仅在训练中被初始化一次,则无需指定stream,会自动创建一个stream。
26- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-Swish.md+1-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.Swish()22torch_npu.contrib.Swish()
23```23```
24 24 
@@ -31,4 +31,3 @@ torch_npu.contrib.Swish()
31>>> input_tensor = torch.randn(2, 32, 5, 5).npu()31>>> input_tensor = torch.randn(2, 32, 5, 5).npu()
32>>> output = m(input_tensor)32>>> output = m(input_tensor)
33```33```
34- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-dropout_with_byte_mask.md+4-2
@@ -15,11 +15,12 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.function.dropout_with_byte_mask(input1, p=0.5, training=True, inplace=False)19torch_npu.contrib.function.dropout_with_byte_mask(input1, p=0.5, training=True, inplace=False)
20```20```
21 21 
22## 参数说明22## 参数说明
23+ 
23- **input1** (`Tensor`): 必选参数,输入张量。24- **input1** (`Tensor`): 必选参数,输入张量。
24- **p** (`float`):可选参数,dropout概率,默认值为0.5。25- **p** (`float`):可选参数,dropout概率,默认值为0.5。
25- **training** (`bool`):可选参数,是否启动dropout,当设置为True时启动,False时不启动。默认值为True。26- **training** (`bool`):可选参数,是否启动dropout,当设置为True时启动,False时不启动。默认值为True。
@@ -30,6 +31,7 @@ torch_npu.contrib.function.dropout_with_byte_mask(input1, p=0.5, training=True,
30仅在设备32核场景下性能提升。31仅在设备32核场景下性能提升。
31 32 
32## 使用示例33## 使用示例
34+ 
33```python35```python
34import torch, torch_npu36import torch, torch_npu
35from torch_npu.contrib.function import npu_functional as F37from torch_npu.contrib.function import npu_functional as F
@@ -37,4 +39,4 @@ input = torch.randn(4,4).npu()
37input = torch_npu.npu_format_cast(input, 2)39input = torch_npu.npu_format_cast(input, 2)
38output = F.dropout_with_byte_mask(input, p=0.2, training=True)40output = F.dropout_with_byte_mask(input, p=0.2, training=True)
39output41output
40-```42+```
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-fuse_add_softmax_dropout.md+4-1
@@ -16,6 +16,7 @@
16* 等价计算逻辑:16* 等价计算逻辑:
17 17 
18 可使用`npu_fuse_add_softmax_dropout`等价替换`torch_npu.contrib.function.fuse_add_softmax_dropout`,两者计算逻辑一致。18 可使用`npu_fuse_add_softmax_dropout`等价替换`torch_npu.contrib.function.fuse_add_softmax_dropout`,两者计算逻辑一致。
19+ 
19 ```python20 ```python
20 import torch21 import torch
21 import math22 import math
@@ -26,9 +27,10 @@
26 attn_probs = dropout(attn_probs)27 attn_probs = dropout(attn_probs)
27 return attn_probs28 return attn_probs
28 ```29 ```
30+ 
29## 函数原型31## 函数原型
30 32 
31-```33+```python
32torch_npu.contrib.function.fuse_add_softmax_dropout(training, dropout, attn_mask, attn_scores, attn_head_size, p=0.5, dim=-1) -> Tensor34torch_npu.contrib.function.fuse_add_softmax_dropout(training, dropout, attn_mask, attn_scores, attn_head_size, p=0.5, dim=-1) -> Tensor
33```35```
34 36 
@@ -43,6 +45,7 @@ torch_npu.contrib.function.fuse_add_softmax_dropout(training, dropout, attn_mask
43- **dim** (`int`):可选参数,待计算softmax的维度,默认值为-1。45- **dim** (`int`):可选参数,待计算softmax的维度,默认值为-1。
44 46 
45## 返回值说明47## 返回值说明
48+ 
46`Tensor`49`Tensor`
47 50 
48返回计算结果。51返回计算结果。
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-matmul_transpose.md+1-3
@@ -9,14 +9,13 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15使用NPU自定义算子替换原生写法,以提高性能。14使用NPU自定义算子替换原生写法,以提高性能。
16 15 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.contrib.function.matmul_transpose(tensor1, tensor2)19torch_npu.contrib.function.matmul_transpose(tensor1, tensor2)
21```20```
22 21 
@@ -50,4 +49,3 @@ torch_npu.contrib.function.matmul_transpose(tensor1, tensor2)
50>>> output.shape49>>> output.shape
51torch.Size([68, 5, 75, 75])50torch.Size([68, 5, 75, 75])
52```51```
53- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-npu_batched_multiclass_nms.md+2-6
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.function.npu_batched_multiclass_nms1# (beta)torch_npu.contrib.function.npu_batched_multiclass_nms
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -15,12 +15,10 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.function.npu_batched_multiclass_nms(multi_bboxes, multi_scores, score_thr=0.05, nms_thr=0.45, max_num=50, score_factors=None)19torch_npu.contrib.function.npu_batched_multiclass_nms(multi_bboxes, multi_scores, score_thr=0.05, nms_thr=0.45, max_num=50, score_factors=None)
20```20```
21 21 
22- 
23- 
24## 参数说明22## 参数说明
25 23 
26- **multi_bboxes** (`Tensor`): 必选参数。候选框(bbox)张量,shape为(bs, n, class, 4)或(bs, n, 4)。24- **multi_bboxes** (`Tensor`): 必选参数。候选框(bbox)张量,shape为(bs, n, class, 4)或(bs, n, 4)。
@@ -40,7 +38,6 @@ torch_npu.contrib.function.npu_batched_multiclass_nms(multi_bboxes, multi_scores
40 38 
41在动态shape条件下,最多支持20个类别(nmsed_classes)和10000个框(nmsed_boxes)。39在动态shape条件下,最多支持20个类别(nmsed_classes)和10000个框(nmsed_boxes)。
42 40 
43- 
44## 调用示例41## 调用示例
45 42 
46```python43```python
@@ -54,4 +51,3 @@ torch.Size([4, 3, 5])
54>>> det_labels.shape51>>> det_labels.shape
55torch.Size([4, 3])52torch.Size([4, 3])
56```53```
57- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-npu_bbox_coder_decode_xywh2xyxy.md+1-2
@@ -15,7 +15,7 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.function.npu_bbox_coder_decode_xywh2xyxy(bboxes, pred_bboxes, means=None, stds=None, max_shape=[9999, 9999], wh_ratio_clip=16 / 1000)19torch_npu.contrib.function.npu_bbox_coder_decode_xywh2xyxy(bboxes, pred_bboxes, means=None, stds=None, max_shape=[9999, 9999], wh_ratio_clip=16 / 1000)
20```20```
21 21 
@@ -48,4 +48,3 @@ torch_npu.contrib.function.npu_bbox_coder_decode_xywh2xyxy(bboxes, pred_bboxes,
48>>> print('npu_bbox_coder_decode_xywh2xyxy done. output shape is ', out.shape)48>>> print('npu_bbox_coder_decode_xywh2xyxy done. output shape is ', out.shape)
49npu_bbox_coder_decode_xywh2xyxy done. output shape is torch.Size([1024, 4])49npu_bbox_coder_decode_xywh2xyxy done. output shape is torch.Size([1024, 4])
50```50```
51- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-npu_bbox_coder_encode_xyxy2xywh.md+2-2
@@ -15,7 +15,7 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.function.npu_bbox_coder_encode_xyxy2xywh(bboxes,gt_bboxes, means=None, stds=None, is_normalized=False, normalized_scale=10000.)19torch_npu.contrib.function.npu_bbox_coder_encode_xyxy2xywh(bboxes,gt_bboxes, means=None, stds=None, is_normalized=False, normalized_scale=10000.)
20```20```
21 21 
@@ -29,6 +29,7 @@ torch_npu.contrib.function.npu_bbox_coder_encode_xyxy2xywh(bboxes,gt_bboxes, mea
29- **normalized_scale** (`Float`):设置坐标恢复的归一化比例,默认值为10000.。29- **normalized_scale** (`Float`):设置坐标恢复的归一化比例,默认值为10000.。
30 30 
31## 返回值说明31## 返回值说明
32+ 
32`Tensor`33`Tensor`
33 34 
34代表框转换deltas。35代表框转换deltas。
@@ -50,4 +51,3 @@ torch_npu.contrib.function.npu_bbox_coder_encode_xyxy2xywh(bboxes,gt_bboxes, mea
50>>> print('npu_bbox_coder_encode_xyxy2xywh done. output shape is ', out.shape)51>>> print('npu_bbox_coder_encode_xyxy2xywh done. output shape is ', out.shape)
51npu_bbox_coder_encode_xyxy2xywh done. output shape is torch.Size([1024, 4])52npu_bbox_coder_encode_xyxy2xywh done. output shape is torch.Size([1024, 4])
52```53```
53- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-npu_bbox_coder_encode_yolo.md+4-3
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.function.npu_bbox_coder_encode_yolo1# (beta)torch_npu.contrib.function.npu_bbox_coder_encode_yolo
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -8,13 +8,14 @@
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11+ 
11## 功能说明12## 功能说明
12 13 
13通过 NPU OP来计算从源框(bbox)到目标框(gt_bbox)的 YOLO 风格框回归转换 deltas。14通过 NPU OP来计算从源框(bbox)到目标框(gt_bbox)的 YOLO 风格框回归转换 deltas。
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.contrib.function.npu_bbox_coder_encode_yolo(bboxes, gt_bboxes, stride)19torch_npu.contrib.function.npu_bbox_coder_encode_yolo(bboxes, gt_bboxes, stride)
19```20```
20 21 
@@ -25,6 +26,7 @@ torch_npu.contrib.function.npu_bbox_coder_encode_yolo(bboxes, gt_bboxes, stride)
25- **stride** (`Tensor`):bbox步长。仅支持`int`张量。26- **stride** (`Tensor`):bbox步长。仅支持`int`张量。
26 27 
27## 返回值说明28## 返回值说明
29+ 
28`Tensor`30`Tensor`
29 31 
30框转换deltas。32框转换deltas。
@@ -43,4 +45,3 @@ torch_npu.contrib.function.npu_bbox_coder_encode_yolo(bboxes, gt_bboxes, stride)
43>>> print('npu_bbox_coder_encode_yolo done. output shape is ', out.shape)45>>> print('npu_bbox_coder_encode_yolo done. output shape is ', out.shape)
44npu_bbox_coder_encode_yolo done. output shape is torch.Size([1024, 4])46npu_bbox_coder_encode_yolo done. output shape is torch.Size([1024, 4])
45```47```
46- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-npu_ciou.md+1-3
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.function.npu_ciou(boxes1, boxes2, trans=True, is_cross=False, mode=0)22torch_npu.contrib.function.npu_ciou(boxes1, boxes2, trans=True, is_cross=False, mode=0)
23```23```
24 24 
@@ -40,7 +40,6 @@ torch_npu.contrib.function.npu_ciou(boxes1, boxes2, trans=True, is_cross=False,
40 40 
41到目前为止,CIoU向后只支持当前版本中的trans==True、is_cross==False、mode==0('iou')。如果需要反向传播,确保参数正确。41到目前为止,CIoU向后只支持当前版本中的trans==True、is_cross==False、mode==0('iou')。如果需要反向传播,确保参数正确。
42 42 
43- 
44## 调用示例43## 调用示例
45 44 
46```python45```python
@@ -53,4 +52,3 @@ torch_npu.contrib.function.npu_ciou(boxes1, boxes2, trans=True, is_cross=False,
53>>> l = ciou.sum()52>>> l = ciou.sum()
54>>> l.backward()53>>> l.backward()
55```54```
56- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-npu_diou.md+1-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.function.npu_diou(boxes1, boxes2, trans=True, is_cross=False, mode=0)22torch_npu.contrib.function.npu_diou(boxes1, boxes2, trans=True, is_cross=False, mode=0)
23```23```
24 24 
@@ -52,4 +52,3 @@ torch_npu.contrib.function.npu_diou(boxes1, boxes2, trans=True, is_cross=False,
52>>> l = diou.sum()52>>> l = diou.sum()
53>>> l.backward()53>>> l.backward()
54```54```
55- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-npu_fast_condition_index_put.md+1-3
@@ -9,14 +9,13 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15使用NPU亲和写法替换bool型index_put函数中的原生写法。14使用NPU亲和写法替换bool型index_put函数中的原生写法。
16 15 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.contrib.function.npu_fast_condition_index_put(x, condition, value)19torch_npu.contrib.function.npu_fast_condition_index_put(x, condition, value)
21```20```
22 21 
@@ -56,4 +55,3 @@ tensor([[0.9661, 1.6750, 0.0000, ..., 0.0000, 0.0000, 0.0000],
56torch.Size([128, 8192])55torch.Size([128, 8192])
57 56 
58```57```
59- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-npu_multiclass_nms.md+2-4
@@ -1,9 +1,7 @@
1# (beta)torch_npu.contrib.function.npu_multiclass_nms1# (beta)torch_npu.contrib.function.npu_multiclass_nms
2 2 
3- 
4## 产品支持情况3## 产品支持情况
5 4 
6- 
7| 产品 | 是否支持 |5| 产品 | 是否支持 |
8| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
9|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
@@ -17,7 +15,7 @@
17 15 
18## 函数原型16## 函数原型
19 17 
20-```18+```python
21torch_npu.contrib.function.npu_multiclass_nms(multi_bboxes, multi_scores, score_thr=0.05, nms_thr=0.45, max_num=50, score_factors=None)19torch_npu.contrib.function.npu_multiclass_nms(multi_bboxes, multi_scores, score_thr=0.05, nms_thr=0.45, max_num=50, score_factors=None)
22```20```
23 21 
@@ -31,6 +29,7 @@ torch_npu.contrib.function.npu_multiclass_nms(multi_bboxes, multi_scores, score_
31- **score_factors** (`Tensor`): 可选参数,默认值为None。NMS应用前用来乘分数的因子。29- **score_factors** (`Tensor`): 可选参数,默认值为None。NMS应用前用来乘分数的因子。
32 30 
33## 返回值说明31## 返回值说明
32+ 
34`Tuple`33`Tuple`
35 34 
36表示候选框和标签(bboxes, labels),shape为(k, 5)和(k)的张量。标签以0为基础。35表示候选框和标签(bboxes, labels),shape为(k, 5)和(k)的张量。标签以0为基础。
@@ -52,4 +51,3 @@ torch.Size([3, 5])
52>>> det_labels.shape51>>> det_labels.shape
53torch.Size([3])52torch.Size([3])
54```53```
55- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-npu_single_level_responsible_flags.md+1-3
@@ -15,7 +15,7 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.function.npu_single_level_responsible_flags(featmap_size, gt_bboxes, stride, num_base_anchors)19torch_npu.contrib.function.npu_single_level_responsible_flags(featmap_size, gt_bboxes, stride, num_base_anchors)
20```20```
21 21 
@@ -32,7 +32,6 @@ torch_npu.contrib.function.npu_single_level_responsible_flags(featmap_size, gt_b
32 32 
33代表单层特征图中每个锚点的有效标志。输出大小为[featmap_size[0] \* featmap_size[1] \* num_base_anchors]。33代表单层特征图中每个锚点的有效标志。输出大小为[featmap_size[0] \* featmap_size[1] \* num_base_anchors]。
34 34 
35- 
36## 调用示例35## 调用示例
37 36 
38```python37```python
@@ -52,4 +51,3 @@ torch.Size([1200]) tensor(1, device='npu:0', dtype=torch.uint8) tensor(0, device
52torch.Size([4800]) tensor(1, device='npu:0', dtype=torch.uint8) tensor(0, device='npu:0', dtype=torch.uint8)51torch.Size([4800]) tensor(1, device='npu:0', dtype=torch.uint8) tensor(0, device='npu:0', dtype=torch.uint8)
53 52 
54```53```
55- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-function-roll.md+2-4
@@ -2,7 +2,6 @@
2 2 
3## 产品支持情况3## 产品支持情况
4 4 
5- 
6| 产品 | 是否支持 |5| 产品 | 是否支持 |
7| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
8|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
@@ -15,11 +14,11 @@
15使用NPU亲和写法替换swin-transformer中的原生roll。14使用NPU亲和写法替换swin-transformer中的原生roll。
16 15 
17## 函数原型16## 函数原型
18-```17+ 
18+```python
19torch_npu.contrib.function.roll(input1, shifts, dims)19torch_npu.contrib.function.roll(input1, shifts, dims)
20```20```
21 21 
22- 
23## 参数说明22## 参数说明
24 23 
25- **input1** (`Tensor`):输入张量。24- **input1** (`Tensor`):输入张量。
@@ -47,4 +46,3 @@ torch_npu.contrib.function.roll(input1, shifts, dims)
47>>> shifted_x_npu.shape46>>> shifted_x_npu.shape
48torch.Size([32, 56, 56, 16])47torch.Size([32, 56, 56, 16])
49```48```
50- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-ChannelShuffle.md+5-5
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.module.ChannelShuffle1# (beta)torch_npu.contrib.module.ChannelShuffle
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -16,6 +16,7 @@
16- 等价计算逻辑:16- 等价计算逻辑:
17 17
18 split_shuffle=False场景可使用`cpu_channel_shuffle`等价替换`torch_npu.contrib.module.ChannelShuffle`,两者计算逻辑一致。18 split_shuffle=False场景可使用`cpu_channel_shuffle`等价替换`torch_npu.contrib.module.ChannelShuffle`,两者计算逻辑一致。
19+ 
19 ```python20 ```python
20 import torch21 import torch
21 def cpu_channel_shuffle(x, groups, split_shuffle):22 def cpu_channel_shuffle(x, groups, split_shuffle):
@@ -36,19 +37,20 @@
36 37
37## 函数原型38## 函数原型
38 39 
39-```40+```python
40torch_npu.contrib.module.ChannelShuffle(in_channels, groups=2, split_shuffle=True)41torch_npu.contrib.module.ChannelShuffle(in_channels, groups=2, split_shuffle=True)
41```42```
42 43 
43## 参数说明44## 参数说明
45+ 
44**计算参数**46**计算参数**
45 47 
46- **in_channels** (`int`):必选参数。输入张量中的通道总数。48- **in_channels** (`int`):必选参数。输入张量中的通道总数。
47- **groups** (`int`):可选参数。shuffle组数。默认值为2。49- **groups** (`int`):可选参数。shuffle组数。默认值为2。
48- **split_shuffle** (`bool`):可选参数。shuffle后是否执行chunk操作。默认值为True。50- **split_shuffle** (`bool`):可选参数。shuffle后是否执行chunk操作。默认值为True。
49 51 
50- 
51**计算输入**52**计算输入**
53+ 
52- **x1** (`Tensor`):输入张量。 shape为$(N, C_{in}, *)$。54- **x1** (`Tensor`):输入张量。 shape为$(N, C_{in}, *)$。
53- **x2** (`Tensor`):输入张量。 shape为$(N, C_{in}, *)$。55- **x2** (`Tensor`):输入张量。 shape为$(N, C_{in}, *)$。
54 56 
@@ -61,7 +63,6 @@ torch_npu.contrib.module.ChannelShuffle(in_channels, groups=2, split_shuffle=Tru
61 63 
62只实现了groups为2场景,请自行修改其他groups场景。64只实现了groups为2场景,请自行修改其他groups场景。
63 65 
64- 
65## 调用示例66## 调用示例
66 67 
67```python68```python
@@ -76,4 +77,3 @@ torch.Size([2, 32, 7, 7])
76>>> out2.shape77>>> out2.shape
77torch.Size([2, 32, 7, 7])78torch.Size([2, 32, 7, 7])
78```79```
79- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-Focus.md+3-3
@@ -15,12 +15,12 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.module.Focus(c1, c2, k=1, s=1, p=None, g=1, act=True)19torch_npu.contrib.module.Focus(c1, c2, k=1, s=1, p=None, g=1, act=True)
20```20```
21 21 
22- 
23## 参数说明22## 参数说明
23+ 
24**计算参数**24**计算参数**
25 25 
26- **c1** (`int`):输入图像中的通道数。26- **c1** (`int`):输入图像中的通道数。
@@ -36,6 +36,7 @@ torch_npu.contrib.module.Focus(c1, c2, k=1, s=1, p=None, g=1, act=True)
36- **x**(`Tensor`): 输入张量。36- **x**(`Tensor`): 输入张量。
37 37 
38## 返回值说明38## 返回值说明
39+ 
39`Tensor`40`Tensor`
40 41 
41Focus计算结果。42Focus计算结果。
@@ -53,4 +54,3 @@ Focus计算结果。
53>>> output.shape54>>> output.shape
54torch.Size([4, 13, 150, 20])55torch.Size([4, 13, 150, 20])
55```56```
56- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-FusedColorJitter.md+1-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.module.FusedColorJitter(torch.nn.Module)22torch_npu.contrib.module.FusedColorJitter(torch.nn.Module)
23```23```
24 24 
@@ -40,4 +40,3 @@ torch_npu.contrib.module.FusedColorJitter(torch.nn.Module)
40>>> fcj = FusedColorJitter(0.1, 0.1, 0.1, 0.1).npu()40>>> fcj = FusedColorJitter(0.1, 0.1, 0.1, 0.1).npu()
41>>> img = fcj(image)41>>> img = fcj(image)
42```42```
43- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-LabelSmoothingCrossEntropy.md+4-4
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.module.LabelSmoothingCrossEntropy1# (beta)torch_npu.contrib.module.LabelSmoothingCrossEntropy
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -8,17 +8,19 @@
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11+ 
11## 功能说明12## 功能说明
12 13 
13使用NPU API进行LabelSmoothing Cross Entropy。14使用NPU API进行LabelSmoothing Cross Entropy。
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.contrib.module.LabelSmoothingCrossEntropy(num_classes=1000, smooth_factor=0.)19torch_npu.contrib.module.LabelSmoothingCrossEntropy(num_classes=1000, smooth_factor=0.)
19```20```
20 21 
21## 参数说明22## 参数说明
23+ 
22**计算参数**24**计算参数**
23 25 
24- **num_classes** (`float`):用于onehot的class数量。26- **num_classes** (`float`):用于onehot的class数量。
@@ -35,7 +37,6 @@ torch_npu.contrib.module.LabelSmoothingCrossEntropy(num_classes=1000, smooth_fac
35 37 
36交叉熵计算结果。38交叉熵计算结果。
37 39 
38- 
39## 调用示例40## 调用示例
40 41 
41```python42```python
@@ -50,4 +51,3 @@ torch_npu.contrib.module.LabelSmoothingCrossEntropy(num_classes=1000, smooth_fac
50>>> npu_output51>>> npu_output
51tensor(1.9443, device='npu:0', grad_fn=<MeanBackward1>)52tensor(1.9443, device='npu:0', grad_fn=<MeanBackward1>)
52```53```
53- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-Mish.md+1-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.module.Mish(nn.Module)22torch_npu.contrib.module.Mish(nn.Module)
23```23```
24 24 
@@ -31,4 +31,3 @@ torch_npu.contrib.module.Mish(nn.Module)
31>>> input_tensor = torch.randn(2, 32, 5, 5).npu()31>>> input_tensor = torch.randn(2, 32, 5, 5).npu()
32>>> output = m(input_tensor)32>>> output = m(input_tensor)
33```33```
34- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-ModulatedDeformConv.md+4-4
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.module.ModulatedDeformConv1# (beta)torch_npu.contrib.module.ModulatedDeformConv
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -15,11 +15,12 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.module.ModulatedDeformConv(in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, deformable_groups=1, bias=True, pack=True)19torch_npu.contrib.module.ModulatedDeformConv(in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, deformable_groups=1, bias=True, pack=True)
20```20```
21 21 
22## 参数说明22## 参数说明
23+ 
23**计算参数**24**计算参数**
24 25 
25- **in_channels** (`int`):输入图像中的通道数。26- **in_channels** (`int`):输入图像中的通道数。
@@ -38,6 +39,7 @@ torch_npu.contrib.module.ModulatedDeformConv(in_channels, out_channels, kernel_s
38- **x**(`Tensor`): 输入张量。39- **x**(`Tensor`): 输入张量。
39 40 
40## 返回值说明41## 返回值说明
42+ 
41`Tensor`43`Tensor`
42 44 
43卷积计算结果。45卷积计算结果。
@@ -46,7 +48,6 @@ torch_npu.contrib.module.ModulatedDeformConv(in_channels, out_channels, kernel_s
46 48 
47ModulatedDeformConv仅实现float32数据类型的操作。conv_offset中权重和偏置必须初始化为0。49ModulatedDeformConv仅实现float32数据类型的操作。conv_offset中权重和偏置必须初始化为0。
48 50 
49- 
50## 调用示例51## 调用示例
51 52 
52```python53```python
@@ -58,4 +59,3 @@ ModulatedDeformConv仅实现float32数据类型的操作。conv_offset中权重
58>>> output.shape59>>> output.shape
59torch.Size([2, 32, 5, 5])60torch.Size([2, 32, 5, 5])
60```61```
61- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-MultiheadAttention.md+3-3
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.module.MultiheadAttention1# (beta)torch_npu.contrib.module.MultiheadAttention
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -15,11 +15,12 @@ Multi-head attention。
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.module.MultiheadAttention(embed_dim, num_heads, kdim=None, vdim=None, dropout=0.0, bias=True, add_bias_kv=False, add_zero_attn=False, self_attention=False, encoder_decoder_attention=False, q_noise=0.0, qn_block_size=8)19torch_npu.contrib.module.MultiheadAttention(embed_dim, num_heads, kdim=None, vdim=None, dropout=0.0, bias=True, add_bias_kv=False, add_zero_attn=False, self_attention=False, encoder_decoder_attention=False, q_noise=0.0, qn_block_size=8)
20```20```
21 21 
22## 参数说明22## 参数说明
23+ 
23- **embed_dim** (`int`):模型总维度。24- **embed_dim** (`int`):模型总维度。
24- **num_heads** (`int`):并行attention head。25- **num_heads** (`int`):并行attention head。
25- **kdim**(`int`):key的特性总数。默认值为None。26- **kdim**(`int`):key的特性总数。默认值为None。
@@ -68,4 +69,3 @@ Multi-head attention的计算结果。
68 device='npu:0', dtype=torch.float16,69 device='npu:0', dtype=torch.float16,
69 grad_fn=<NpuMultiHeadAttentionBackward0>), None)70 grad_fn=<NpuMultiHeadAttentionBackward0>), None)
70```71```
71- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-NpuCachedDropout.md+4-3
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.module.NpuCachedDropout1# (beta)torch_npu.contrib.module.NpuCachedDropout
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -8,13 +8,14 @@
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11+ 
11## 功能说明12## 功能说明
12 13 
13在NPU设备上使用FairseqDropout。14在NPU设备上使用FairseqDropout。
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.contrib.module.NpuCachedDropout(p, module_name=None)19torch_npu.contrib.module.NpuCachedDropout(p, module_name=None)
19```20```
20 21 
@@ -25,4 +26,4 @@ torch_npu.contrib.module.NpuCachedDropout(p, module_name=None)
25 26 
26## 约束说明27## 约束说明
27 28 
28-不支持动态shape。29+不支持动态shape。
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-NpuDropPath.md+2-2
@@ -15,7 +15,7 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.module.NpuDropPath(drop_prob=None)19torch_npu.contrib.module.NpuDropPath(drop_prob=None)
20```20```
21 21 
@@ -30,6 +30,7 @@ torch_npu.contrib.module.NpuDropPath(drop_prob=None)
30- **x** (`Tensor`):应用dropout的输入张量。30- **x** (`Tensor`):应用dropout的输入张量。
31 31 
32## 返回值说明32## 返回值说明
33+ 
33`Tensor`34`Tensor`
34 35 
35dropout的计算结果。36dropout的计算结果。
@@ -47,4 +48,3 @@ dropout的计算结果。
47>>> output = input1 + fast_drop_path(input2)48>>> output = input1 + fast_drop_path(input2)
48>>> output.sum().backward()49>>> output.sum().backward()
49```50```
50- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-PSROIPool.md+3-5
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.module.PSROIPool1# (beta)torch_npu.contrib.module.PSROIPool
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -8,14 +8,14 @@
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11+ 
11## 功能说明12## 功能说明
12 13 
13使用NPU API进行PSROIPool。14使用NPU API进行PSROIPool。
14 15 
15- 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.module.PSROIPool(nn.Module)19torch_npu.contrib.module.PSROIPool(nn.Module)
20```20```
21 21 
@@ -37,11 +37,9 @@ shape为(k, 5)和(k, 1)的张量。标签以0为基础。
37 37 
38仅实现了pooled_height == pooled_width == group_size。38仅实现了pooled_height == pooled_width == group_size。
39 39 
40- 
41## 调用示例40## 调用示例
42 41 
43```python42```python
44>>> from torch_npu.contrib.module import PSROIPool43>>> from torch_npu.contrib.module import PSROIPool
45>>> model = PSROIPool(pooled_height=7, pooled_width=7, spatial_scale=1 / 16.0, group_size=7, output_dim=22)44>>> model = PSROIPool(pooled_height=7, pooled_width=7, spatial_scale=1 / 16.0, group_size=7, output_dim=22)
46```45```
47- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-ROIAlign.md+2-7
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.module.ROIAlign1# (beta)torch_npu.contrib.module.ROIAlign
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -13,14 +13,12 @@
13 13 
14使用NPU API进行ROIAlign。14使用NPU API进行ROIAlign。
15 15 
16- 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.contrib.module.ROIAlign(output_size, spatial_scale, sampling_ratio, aligned=True)19torch_npu.contrib.module.ROIAlign(output_size, spatial_scale, sampling_ratio, aligned=True)
21```20```
22 21 
23- 
24## 参数说明22## 参数说明
25 23 
26**计算参数**24**计算参数**
@@ -39,14 +37,12 @@ torch_npu.contrib.module.ROIAlign(output_size, spatial_scale, sampling_ratio, al
39- **input_tensor**(`Tensor`): 输入张量,格式为NCHW。37- **input_tensor**(`Tensor`): 输入张量,格式为NCHW。
40- **rois**(`Tensor`): roi框,2D张量,第二个维度size为5,第一列表示roi框的索引,其余4列为roi框的坐标。38- **rois**(`Tensor`): roi框,2D张量,第二个维度size为5,第一列表示roi框的索引,其余4列为roi框的坐标。
41 39 
42- 
43## 返回值说明40## 返回值说明
44 41 
45`Tensor` 42`Tensor`
46 43 
47ROIAlign计算结果。44ROIAlign计算结果。
48 45 
49- 
50## 调用示例46## 调用示例
51 47 
52```python48```python
@@ -66,4 +62,3 @@ ROIAlign计算结果。
66>>> output.shape62>>> output.shape
67torch.Size([1, 1, 3, 3])63torch.Size([1, 1, 3, 3])
68```64```
69- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-SiLU.md+2-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.module.SiLU(nn.Module)22torch_npu.contrib.module.SiLU(nn.Module)
23```23```
24 24 
@@ -31,4 +31,4 @@ torch_npu.contrib.module.SiLU(nn.Module)
31>>> m = SiLU()31>>> m = SiLU()
32>>> input_tensor = torch.randn(2, 32, 5, 5).npu()32>>> input_tensor = torch.randn(2, 32, 5, 5).npu()
33>>> output = m(input_tensor)33>>> output = m(input_tensor)
34-```34+```
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-module-npu_modules-DropoutWithByteMask.md+3-4
@@ -1,6 +1,6 @@
1# (beta)torch_npu.contrib.module.npu_modules.DropoutWithByteMask1# (beta)torch_npu.contrib.module.npu_modules.DropoutWithByteMask
2-## 产品支持情况
3 2 
3+## 产品支持情况
4 4 
5| 产品 | 是否支持 |5| 产品 | 是否支持 |
6| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
@@ -8,13 +8,14 @@
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11+ 
11## 功能说明12## 功能说明
12 13 
13应用NPU兼容的DropoutWithByteMask操作。14应用NPU兼容的DropoutWithByteMask操作。
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.contrib.module.npu_modules.DropoutWithByteMask(p=0.5, inplace=False, max_seed=2 ** 10 - 1)19torch_npu.contrib.module.npu_modules.DropoutWithByteMask(p=0.5, inplace=False, max_seed=2 ** 10 - 1)
19```20```
20 21 
@@ -36,7 +37,6 @@ torch_npu.contrib.module.npu_modules.DropoutWithByteMask(p=0.5, inplace=False, m
36 37 
37输出张量与输入张量的shape相同。38输出张量与输入张量的shape相同。
38 39 
39- 
40## 调用示例40## 调用示例
41 41 
42```python42```python
@@ -48,4 +48,3 @@ torch_npu.contrib.module.npu_modules.DropoutWithByteMask(p=0.5, inplace=False, m
48>>> output.shape48>>> output.shape
49torch.Size([16, 16])49torch.Size([16, 16])
50```50```
51- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-npu_fused_attention.md+1-2
@@ -15,7 +15,7 @@ bert自注意力的融合实现。
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.contrib.npu_fused_attention(hidden_states, attention_mask, query_kernel, key_kernel, value_kernel, query_bias, key_bias, value_bias, scale=1, keep_prob=0)19torch_npu.contrib.npu_fused_attention(hidden_states, attention_mask, query_kernel, key_kernel, value_kernel, query_bias, key_bias, value_bias, scale=1, keep_prob=0)
20```20```
21 21 
@@ -37,4 +37,3 @@ torch_npu.contrib.npu_fused_attention(hidden_states, attention_mask, query_kerne
37`Tensor`37`Tensor`
38 38 
39self attention的结果。39self attention的结果。
40- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-npu_fused_attention_with_layernorm.md+1-2
@@ -18,7 +18,7 @@ bert自注意力与层归一化的融合实现。
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.npu_fused_attention_with_layernorm(hidden_states, attention_mask, query_kernel, key_kernel, value_kernel, query_bias, key_bias, value_bias, gamma, beta, scale=1, keep_prob=0)22torch_npu.contrib.npu_fused_attention_with_layernorm(hidden_states, attention_mask, query_kernel, key_kernel, value_kernel, query_bias, key_bias, value_bias, gamma, beta, scale=1, keep_prob=0)
23```23```
24 24 
@@ -42,4 +42,3 @@ torch_npu.contrib.npu_fused_attention_with_layernorm(hidden_states, attention_ma
42`torch.Tensor`42`torch.Tensor`
43 43 
44self attention的结果。44self attention的结果。
45- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-npu_giou.md+1-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.npu_giou(boxes1, boxes2, is_permuted=True)22torch_npu.contrib.npu_giou(boxes1, boxes2, is_permuted=True)
23```23```
24 24 
@@ -40,4 +40,3 @@ torch_npu.contrib.npu_giou(boxes1, boxes2, is_permuted=True)
40>>> l = iou.sum()40>>> l = iou.sum()
41>>> l.backward()41>>> l.backward()
42```42```
43- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-npu_iou.md+1-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.npu_iou(boxes1, boxes2, mode="ptiou", is_normalized=False, normalized_scale=100.)22torch_npu.contrib.npu_iou(boxes1, boxes2, mode="ptiou", is_normalized=False, normalized_scale=100.)
23```23```
24 24 
@@ -40,4 +40,3 @@ torch_npu.contrib.npu_iou(boxes1, boxes2, mode="ptiou", is_normalized=False, nor
40>>> box2 = torch.randint(0, 256, size=(16, 4)).npu()40>>> box2 = torch.randint(0, 256, size=(16, 4)).npu()
41>>> iou = torch_npu.contrib.npu_iou(box1, box2)41>>> iou = torch_npu.contrib.npu_iou(box1, box2)
42```42```
43- 
Mdocs/zh/custom_APIs/torch_npu-contrib/(beta)torch_npu-contrib-npu_ptiou.md+1-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.contrib.npu_ptiou(boxes1, boxes2, mode="ptiou", is_normalized=False, normalized_scale=100.)22torch_npu.contrib.npu_ptiou(boxes1, boxes2, mode="ptiou", is_normalized=False, normalized_scale=100.)
23```23```
24 24 
@@ -42,4 +42,3 @@ torch_npu.contrib.npu_ptiou(boxes1, boxes2, mode="ptiou", is_normalized=False, n
42>>> box2 = torch.randint(0, 256, size=(16, 4)).npu()42>>> box2 = torch.randint(0, 256, size=(16, 4)).npu()
43>>> iou = torch_npu.contrib.npu_ptiou(box1, box2)43>>> iou = torch_npu.contrib.npu_ptiou(box1, box2)
44```44```
45- 
Mdocs/zh/custom_APIs/torch_npu-jit/torch_npu-jit.md+1-1
@@ -1 +1 @@
1-# torch_npu.jit1+# torch_npu.jit
Mdocs/zh/custom_APIs/torch_npu-jit/torch_npu-jit_list.md+0-1
@@ -18,4 +18,3 @@
18</tr>18</tr>
19</tbody>19</tbody>
20</table>20</table>
21- 
Mdocs/zh/custom_APIs/torch_npu-jit/(beta)torch_npu-jit-optimize.md+3-4
@@ -9,13 +9,13 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15主要用于优化ScriptFunction或ScriptModule,以获取更好的性能。14主要用于优化ScriptFunction或ScriptModule,以获取更好的性能。
15+ 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.jit.optimize(jit_mod)19torch_npu.jit.optimize(jit_mod)
20```20```
21 21 
@@ -23,7 +23,6 @@ torch_npu.jit.optimize(jit_mod)
23 23 
24**jit_mod**:必选参数。用于被优化的ScriptFunction或ScriptModule。24**jit_mod**:必选参数。用于被优化的ScriptFunction或ScriptModule。
25 25 
26- 
27## 调用示例26## 调用示例
28 27 
29```python28```python
@@ -41,4 +40,4 @@ traced_model = torch.jit.trace(model, (torch.rand(1, 3), torch.rand(1, 3)))
41 40 
42torch_npu.jit.optimize(traced_model)41torch_npu.jit.optimize(traced_model)
43 42 
44-```43+```
Mdocs/zh/custom_APIs/torch_npu-npu/Memory-management-API.md+0-1
@@ -87,4 +87,3 @@
87</tr>87</tr>
88</tbody>88</tbody>
89</table>89</table>
90- 
Mdocs/zh/custom_APIs/torch_npu-npu/Memory-management.md+1-1
@@ -1 +1 @@
1-# Memory management1+# Memory management
Mdocs/zh/custom_APIs/torch_npu-npu/NPU-Device.md+0-1
@@ -51,4 +51,3 @@
51</tr>51</tr>
52</tbody>52</tbody>
53</table>53</table>
54- 
Mdocs/zh/custom_APIs/torch_npu-npu/Random-Number-Generator.md+0-1
@@ -40,4 +40,3 @@
40</tr>40</tr>
41</tbody>41</tbody>
42</table>42</table>
43- 
Mdocs/zh/custom_APIs/torch_npu-npu/amp.md+1-1
@@ -1 +1 @@
1-# amp1+# amp
Mdocs/zh/custom_APIs/torch_npu-npu/aoe.md+1-1
@@ -1 +1 @@
1-# aoe1+# aoe
Mdocs/zh/custom_APIs/torch_npu-npu/profiler.md+1-1
@@ -1 +1 @@
1-# profiler1+# profiler
Mdocs/zh/custom_APIs/torch_npu-npu/torch-npu-npu-NPUPluggableAllocator.md+2-6
@@ -9,22 +9,20 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13- 
14## 功能说明12## 功能说明
15 13 
16从so文件加载的NPU内存分配器。14从so文件加载的NPU内存分配器。
15+ 
17## 定义文件16## 定义文件
18 17 
19torch_npu/npu/memory.py18torch_npu/npu/memory.py
20 19 
21## 函数原型20## 函数原型
22 21 
23-```22+```python
24torch_npu.npu.NPUPluggableAllocator(path_to_so_file, alloc_fn_name, free_fn_name)23torch_npu.npu.NPUPluggableAllocator(path_to_so_file, alloc_fn_name, free_fn_name)
25```24```
26 25 
27- 
28## 参数说明26## 参数说明
29 27 
30- **path_to_so_file**(`str`):so文件路径。28- **path_to_so_file**(`str`):so文件路径。
@@ -37,7 +35,6 @@ torch_npu.npu.NPUPluggableAllocator(path_to_so_file, alloc_fn_name, free_fn_name
37 35 
38`free_fn_name`内存释放函数名必须与c/c++文件中函数名一致。36`free_fn_name`内存释放函数名必须与c/c++文件中函数名一致。
39 37 
40- 
41## 调用示例38## 调用示例
42 39 
43完整调用示例可参考[LINK](https://gitcode.com/ascend/pytorch/blob/v2.7.1-7.3.0/test/allocator/test_pluggable_allocator_extensions.py)。40完整调用示例可参考[LINK](https://gitcode.com/ascend/pytorch/blob/v2.7.1-7.3.0/test/allocator/test_pluggable_allocator_extensions.py)。
@@ -98,4 +95,3 @@ ASCEND_LOGD("Pluggable Allocator malloc: malloc = %zu", size);
98 95
99ASCEND_LOGD("Pluggable Allocator free: free= %zu", size);96ASCEND_LOGD("Pluggable Allocator free: free= %zu", size);
100```97```
101- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch-npu-npu-change_current_allocator.md+1-9
@@ -9,9 +9,6 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13- 
14- 
15## 功能说明12## 功能说明
16 13 
17更改使用的内存分配器。14更改使用的内存分配器。
@@ -22,12 +19,10 @@ torch_npu/npu/memory.py
22 19 
23## 函数原型20## 函数原型
24 21 
25-```22+```python
26torch_npu.npu.change_current_allocator(allocator) -> None23torch_npu.npu.change_current_allocator(allocator) -> None
27```24```
28 25 
29- 
30- 
31## 参数说明26## 参数说明
32 27 
33 **allocator** (`torch_npu.npu.memory._NPUAllocator`):要设置为使用的内存分配器。28 **allocator** (`torch_npu.npu.memory._NPUAllocator`):要设置为使用的内存分配器。
@@ -36,8 +31,6 @@ torch_npu.npu.change_current_allocator(allocator) -> None
36 31 
37如果内存分配器已被初始化,调用该函数将失败。32如果内存分配器已被初始化,调用该函数将失败。
38 33 
39- 
40- 
41## 调用示例34## 调用示例
42 35 
43完整调用示例可参考[LINK](https://gitcode.com/ascend/pytorch/blob/v2.7.1-7.3.0/test/allocator/test_pluggable_allocator_extensions.py)。36完整调用示例可参考[LINK](https://gitcode.com/ascend/pytorch/blob/v2.7.1-7.3.0/test/allocator/test_pluggable_allocator_extensions.py)。
@@ -96,4 +89,3 @@ ASCEND_LOGD("Pluggable Allocator malloc: malloc = %zu", size);
96 89
97ASCEND_LOGD("Pluggable Allocator free: free= %zu", size);90ASCEND_LOGD("Pluggable Allocator free: free= %zu", size);
98```91```
99- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-Storage.md+0-1
@@ -40,4 +40,3 @@
40</tr>40</tr>
41</tbody>41</tbody>
42</table>42</table>
43- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-1.md+1-1
@@ -1 +1 @@
1-# torch_npu.npu1+# torch_npu.npu
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-2.md+1-1
@@ -1 +1 @@
1-# torch_npu.npu1+# torch_npu.npu
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-Event()-recorded_time-().md+3-4
@@ -13,7 +13,7 @@
13 13 
14## 函数原型14## 函数原型
15 15 
16-```16+```python
17torch_npu.npu.Event().recorded_time() -> int17torch_npu.npu.Event().recorded_time() -> int
18```18```
19 19 
@@ -23,10 +23,10 @@ torch_npu.npu.Event().recorded_time() -> int
23 23 
24## 返回值说明24## 返回值说明
25 25 
26+- **int**:输出被记录的时间,是一个无符号的整数(uint64),单位为微秒。
26 27 
27-- **int**:输出被记录时间,是一个无符号的整数(uint64),单位为微秒28+- 若返回“INTERNALError”,则表示Event对象必须在获取记录时间戳之前被记录
28 29 
29-- 若返回“INTERNALError”,则表示Event对象必须在获取记录时间戳之前被记录。
30## 约束说明30## 约束说明
31 31 
32Event对象在创建的时候,需要传入参数enable_timing=True。32Event对象在创建的时候,需要传入参数enable_timing=True。
@@ -41,4 +41,3 @@ event = torch_npu.npu.Event(enable_timing=True)
41event.record()41event.record()
42res = event.recorded_time()42res = event.recorded_time()
43```43```
44- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-ExternalEvent().record().md+3-2
@@ -1,4 +1,5 @@
1# torch_npu.npu.ExternalEvent().record()1# torch_npu.npu.ExternalEvent().record()
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -10,14 +11,13 @@
10|<term>Atlas 推理系列产品</term> | √ |11|<term>Atlas 推理系列产品</term> | √ |
11|<term>Atlas 训练系列产品</term> | √ |12|<term>Atlas 训练系列产品</term> | √ |
12 13 
13- 
14## 功能说明14## 功能说明
15 15 
16在指定stream上记录Event事件。本接口被调用时,会捕获当前Stream上已下发的任务,并记录到Event事件中,因此后续若调用wait接口,会等待该Event事件中所捕获的任务都已经完成。16在指定stream上记录Event事件。本接口被调用时,会捕获当前Stream上已下发的任务,并记录到Event事件中,因此后续若调用wait接口,会等待该Event事件中所捕获的任务都已经完成。
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.ExternalEvent().record(stream) -> None21torch_npu.npu.ExternalEvent().record(stream) -> None
22```22```
23 23 
@@ -35,6 +35,7 @@ torch_npu.npu.ExternalEvent().record(stream) -> None
35- 接口调用顺序:torch_npu.npu.ExternalEvent().wait()-->torch_npu.npu.ExternalEvent().record()或torch_npu.npu.ExternalEvent().record()-->torch_npu.npu.ExternalEvent().wait()。35- 接口调用顺序:torch_npu.npu.ExternalEvent().wait()-->torch_npu.npu.ExternalEvent().record()或torch_npu.npu.ExternalEvent().record()-->torch_npu.npu.ExternalEvent().wait()。
36 36 
37## 调用示例37## 调用示例
38+ 
38```python39```python
39import torch40import torch
40import torch_npu41import torch_npu
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-ExternalEvent().reset().md+2-2
@@ -14,14 +14,13 @@
14|<term>Atlas 推理系列产品</term> | √ |14|<term>Atlas 推理系列产品</term> | √ |
15|<term>Atlas 训练系列产品</term> | √ |15|<term>Atlas 训练系列产品</term> | √ |
16 16 
17- 
18## 功能说明17## 功能说明
19 18 
20复位一个Event。Event复用场景,用于复位因record任务完成置位的标志位。19复位一个Event。Event复用场景,用于复位因record任务完成置位的标志位。
21 20 
22## 函数原型21## 函数原型
23 22 
24-```23+```python
25torch_npu.npu.ExternalEvent().reset(stream) -> None24torch_npu.npu.ExternalEvent().reset(stream) -> None
26```25```
27 26 
@@ -39,6 +38,7 @@ torch_npu.npu.ExternalEvent().reset(stream) -> None
39- 接口调用顺序:torch_npu.npu.ExternalEvent().wait()-->torch_npu.npu.ExternalEvent().reset()-->torch_npu.npu.ExternalEvent().record()或torch_npu.npu.ExternalEvent().record()-->torch_npu.npu.ExternalEvent().wait()-->torch_npu.npu.ExternalEvent().reset()。38- 接口调用顺序:torch_npu.npu.ExternalEvent().wait()-->torch_npu.npu.ExternalEvent().reset()-->torch_npu.npu.ExternalEvent().record()或torch_npu.npu.ExternalEvent().record()-->torch_npu.npu.ExternalEvent().wait()-->torch_npu.npu.ExternalEvent().reset()。
40 39 
41## 调用示例40## 调用示例
41+ 
42```python42```python
43import torch43import torch
44import torch_npu44import torch_npu
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-ExternalEvent().wait().md+3-2
@@ -1,4 +1,5 @@
1# torch_npu.npu.ExternalEvent().wait()1# torch_npu.npu.ExternalEvent().wait()
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -10,14 +11,13 @@
10|<term>Atlas 推理系列产品</term> | √ |11|<term>Atlas 推理系列产品</term> | √ |
11|<term>Atlas 训练系列产品</term> | √ |12|<term>Atlas 训练系列产品</term> | √ |
12 13 
13- 
14## 功能说明14## 功能说明
15 15 
16阻塞指定Stream的运行,直到指定的Event完成,仅支持单个Stream等待单个Event的场景。16阻塞指定Stream的运行,直到指定的Event完成,仅支持单个Stream等待单个Event的场景。
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.ExternalEvent().wait(stream) -> None21torch_npu.npu.ExternalEvent().wait(stream) -> None
22```22```
23 23 
@@ -36,6 +36,7 @@ torch_npu.npu.ExternalEvent().wait(stream) -> None
36- 该接口会自动复位Event,不需要调用torch_npu.npu.ExternalEvent().reset()接口手动复位Event。36- 该接口会自动复位Event,不需要调用torch_npu.npu.ExternalEvent().reset()接口手动复位Event。
37 37 
38## 调用示例38## 调用示例
39+ 
39```python40```python
40import torch41import torch
41import torch_npu42import torch_npu
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-ExternalEvent.md+4-2
@@ -1,4 +1,5 @@
1# torch_npu.npu.ExternalEvent1# torch_npu.npu.ExternalEvent
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -10,18 +11,18 @@
10|<term>Atlas 推理系列产品</term> | √ |11|<term>Atlas 推理系列产品</term> | √ |
11|<term>Atlas 训练系列产品</term> | √ |12|<term>Atlas 训练系列产品</term> | √ |
12 13 
13- 
14## 功能说明14## 功能说明
15 15 
16ExternalEvent是AscendCL Event的封装。NPUGraph场景在执行图捕获时,ExternalEvent会被作为图外部节点被捕获,用于控制非图内时序控制场景。16ExternalEvent是AscendCL Event的封装。NPUGraph场景在执行图捕获时,ExternalEvent会被作为图外部节点被捕获,用于控制非图内时序控制场景。
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.ExternalEvent()21torch_npu.npu.ExternalEvent()
22```22```
23 23 
24## 返回值说明24## 返回值说明
25+ 
25返回创建好的ExternalEvent对象,用于下发Event相关任务。26返回创建好的ExternalEvent对象,用于下发Event相关任务。
26 27 
27## 约束说明28## 约束说明
@@ -29,6 +30,7 @@ torch_npu.npu.ExternalEvent()
29 ExternalEvent创建时,系统内部会在Device上分配32字节的内存,创建数量受芯片硬件规格限制。30 ExternalEvent创建时,系统内部会在Device上分配32字节的内存,创建数量受芯片硬件规格限制。
30 31 
31## 调用示例32## 调用示例
33+ 
32```python34```python
33import torch35import torch
34import torch_npu36import torch_npu
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-SyncLaunchStream.md+1-3
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.npu.SyncLaunchStream(device)18torch_npu.npu.SyncLaunchStream(device)
19```19```
20 20 
@@ -31,7 +31,6 @@ torch_npu.npu.SyncLaunchStream(device)
31- 由于不再下发到taskqueue,因此该流的下发性能相比普通流有所降低,建议在集群训练时某些节点出现故障,其他节点保存ckpt时创建一条同步下发NPUStream。31- 由于不再下发到taskqueue,因此该流的下发性能相比普通流有所降低,建议在集群训练时某些节点出现故障,其他节点保存ckpt时创建一条同步下发NPUStream。
32- 同步下发流资源池只有4条,创建超过4条时将会循环从资源池中获取。32- 同步下发流资源池只有4条,创建超过4条时将会循环从资源池中获取。
33 33 
34- 
35## 调用示例34## 调用示例
36 35 
37```python36```python
@@ -43,4 +42,3 @@ with torch.npu.stream(s):
43 tensor2 = tensor1 + tensor142 tensor2 = tensor1 + tensor1
44 s.synchronize()43 s.synchronize()
45```44```
46- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-aclnn-allow_hf32.md+1-3
@@ -9,14 +9,13 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15设置或查询conv类算子是否支持hf32。14设置或查询conv类算子是否支持hf32。
16 15 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.npu.aclnn.allow_hf32:bool19torch_npu.npu.aclnn.allow_hf32:bool
21```20```
22 21 
@@ -45,4 +44,3 @@ False
45>>> res44>>> res
46True45True
47```46```
48- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-are_compatible_impl_enabled.md+3-4
@@ -17,7 +17,7 @@
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.are_compatible_impl_enabled()21torch_npu.npu.are_compatible_impl_enabled()
22```22```
23 23 
@@ -26,12 +26,11 @@ torch_npu.npu.are_compatible_impl_enabled()
2626
27 27 
28## 返回值说明28## 返回值说明
29+ 
29`bool`30`bool`
30 31 
31True为已开启,False为未开启。32True为已开启,False为未开启。
32 33 
33- 
34- 
35## 约束说明34## 约束说明
36 35 
3736
@@ -44,4 +43,4 @@ True为已开启,False为未开启。
44>>> torch_npu.npu.use_compatible_impl(True)43>>> torch_npu.npu.use_compatible_impl(True)
45>>> torch_npu.npu.are_compatible_impl_enabled()44>>> torch_npu.npu.are_compatible_impl_enabled()
46True45True
47-```46+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-conv-allow_hf32.md+2-6
@@ -1,4 +1,5 @@
1# torch_npu.npu.conv.allow_hf321# torch_npu.npu.conv.allow_hf32
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -8,8 +9,6 @@
8|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
9|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
10 11 
11- 
12- 
13## 功能说明12## 功能说明
14 13 
15conv类算子开启支持hf32类型能力。14conv类算子开启支持hf32类型能力。
@@ -18,11 +17,10 @@ conv类算子开启支持hf32类型能力。
18 17 
19## 函数原型18## 函数原型
20 19 
21-```20+```python
22torch_npu.npu.conv.allow_hf32 = bool21torch_npu.npu.conv.allow_hf32 = bool
23```22```
24 23 
25- 
26## 参数说明24## 参数说明
27 25 
28输入bool值,默认值True。26输入bool值,默认值True。
@@ -31,7 +29,6 @@ torch_npu.npu.conv.allow_hf32 = bool
31 29 
32`bool`30`bool`
33 31 
34- 
35## 调用示例32## 调用示例
36 33 
37```python34```python
@@ -46,4 +43,3 @@ False
46>>>torch_npu.npu.conv.allow_hf3243>>>torch_npu.npu.conv.allow_hf32
47True44True
48```45```
49- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-disable_deterministic_with_backward.md+2-4
@@ -10,14 +10,13 @@
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11|<term>Atlas 200I/500 A2 推理产品</term> | √ |11|<term>Atlas 200I/500 A2 推理产品</term> | √ |
12 12 
13- 
14## 功能说明13## 功能说明
15 14 
16关闭“确定性”功能。确定性算法是指在模型的前向传播过程中,每次输入相同,输出也相同。15关闭“确定性”功能。确定性算法是指在模型的前向传播过程中,每次输入相同,输出也相同。
17 16 
18## 函数原型17## 函数原型
19 18 
20-```19+```python
21torch_npu.npu.disable_deterministic_with_backward(tensor) -> Tensor20torch_npu.npu.disable_deterministic_with_backward(tensor) -> Tensor
22```21```
23 22 
@@ -26,6 +25,7 @@ torch_npu.npu.disable_deterministic_with_backward(tensor) -> Tensor
26**tensor** (`Tensor`):该接口为透明传输接口,不做数据处理,类型支持和数据格式为PyTorch在各芯片上的可支持的数据类型和数据格式,无接口级别的约束。25**tensor** (`Tensor`):该接口为透明传输接口,不做数据处理,类型支持和数据格式为PyTorch在各芯片上的可支持的数据类型和数据格式,无接口级别的约束。
27 26 
28## 返回值说明27## 返回值说明
28+ 
29`Tensor`29`Tensor`
30 30 
31代表`disable_deterministic_with_backward`的计算结果。31代表`disable_deterministic_with_backward`的计算结果。
@@ -35,7 +35,6 @@ torch_npu.npu.disable_deterministic_with_backward(tensor) -> Tensor
35- 入参`tensor`需要是训练网络中可以传递下去且与整网的`output`有关联的`tensor`变量,否则无法进行反向设置确定性能力。35- 入参`tensor`需要是训练网络中可以传递下去且与整网的`output`有关联的`tensor`变量,否则无法进行反向设置确定性能力。
36- 不支持图模式。36- 不支持图模式。
37 37 
38- 
39## 调用示例38## 调用示例
40 39 
41单算子模式调用:40单算子模式调用:
@@ -91,4 +90,3 @@ Ran 1 test in 4.636s
91 90 
92OK91OK
93```92```
94- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-empty_virt_addr_cache.md+5-5
@@ -11,30 +11,30 @@
11 11 
12轻量化的缓存释放接口,对应于`torch.npu.empty_cache`。只释放虚拟内存,解除虚拟内存与物理内存的映射,但不真正释放物理内存,从而降低调用耗时。12轻量化的缓存释放接口,对应于`torch.npu.empty_cache`。只释放虚拟内存,解除虚拟内存与物理内存的映射,但不真正释放物理内存,从而降低调用耗时。
13 13 
14- 
15## 定义文件14## 定义文件
15+ 
16torch_npu/npu/memory.py16torch_npu/npu/memory.py
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.empty_virt_addr_cache() -> None21torch_npu.npu.empty_virt_addr_cache() -> None
22```22```
23 23 
24## 参数说明24## 参数说明
25+ 
2526
26 27 
27## 返回值说明28## 返回值说明
29+ 
2830
29 31 
30## 约束说明32## 约束说明
31 33 
32该接口需要环境变量`PYTORCH_NPU_ALLOC_CONF`的值设置为`expandable_segments:True`时才生效,否则会runtime报错。34该接口需要环境变量`PYTORCH_NPU_ALLOC_CONF`的值设置为`expandable_segments:True`时才生效,否则会runtime报错。
33 35 
34- 
35## 调用示例36## 调用示例
36 37 
37- 
38```python38```python
39>>> import torch39>>> import torch
40>>> import torch_npu40>>> import torch_npu
@@ -42,4 +42,4 @@ torch_npu.npu.empty_virt_addr_cache() -> None
42>>> del x42>>> del x
43>>> torch_npu.npu.empty_virt_addr_cache()43>>> torch_npu.npu.empty_virt_addr_cache()
44>>> print(torch_npu.npu.memory_summary())44>>> print(torch_npu.npu.memory_summary())
45-```45+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-enable_deterministic_with_backward.md+2-4
@@ -1,6 +1,5 @@
1# torch_npu.npu.enable_deterministic_with_backward1# torch_npu.npu.enable_deterministic_with_backward
2 2 
3- 
4## 产品支持情况3## 产品支持情况
5 4 
6| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -18,7 +17,7 @@
18 17 
19## 函数原型18## 函数原型
20 19 
21-```20+```python
22torch_npu.npu.enable_deterministic_with_backward(tensor) -> Tensor21torch_npu.npu.enable_deterministic_with_backward(tensor) -> Tensor
23```22```
24 23 
@@ -27,6 +26,7 @@ torch_npu.npu.enable_deterministic_with_backward(tensor) -> Tensor
27**tensor** (`Tensor`):该接口为透明传输接口,不做数据处理,类型支持和数据格式为PyTorch在各芯片上可支持的数据类型和数据格式,无接口级别的约束。26**tensor** (`Tensor`):该接口为透明传输接口,不做数据处理,类型支持和数据格式为PyTorch在各芯片上可支持的数据类型和数据格式,无接口级别的约束。
28 27 
29## 返回值说明28## 返回值说明
29+ 
30`Tensor`30`Tensor`
31 31 
32代表`enable_deterministic_with_backward`的计算结果。32代表`enable_deterministic_with_backward`的计算结果。
@@ -36,7 +36,6 @@ torch_npu.npu.enable_deterministic_with_backward(tensor) -> Tensor
36- 入参`tensor`需要是训练网络中可以传递下去且与整网的`output`有关联的`tensor`变量,否则无法进行反向设置确定性能力。36- 入参`tensor`需要是训练网络中可以传递下去且与整网的`output`有关联的`tensor`变量,否则无法进行反向设置确定性能力。
37- 不支持图模式。37- 不支持图模式。
38 38 
39- 
40## 调用示例39## 调用示例
41 40 
42单算子模式调用:41单算子模式调用:
@@ -95,4 +94,3 @@ Ran 1 test in 4.636s
95 94 
96OK95OK
97```96```
98- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-graph_task_group_begin.md+3-3
@@ -1,4 +1,5 @@
1# torch_npu.npu.graph_task_group_begin1# torch_npu.npu.graph_task_group_begin
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -8,14 +9,13 @@
8|<term>Atlas A2 训练系列产品</term> | √ |9|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas A2 推理系列产品</term> | √ |10|<term>Atlas A2 推理系列产品</term> | √ |
10 11 
11- 
12## 功能说明12## 功能说明
13 13 
14NPUGraph场景下,用于标记任务组起始位置。14NPUGraph场景下,用于标记任务组起始位置。
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.npu.graph_task_group_begin(stream) -> None19torch_npu.npu.graph_task_group_begin(stream) -> None
20```20```
21 21 
@@ -32,6 +32,7 @@ torch_npu.npu.graph_task_group_begin(stream) -> None
32图捕获阶段,与[torch_npu.npu.graph_task_group_end](torch_npu-npu-graph_task_group_end.md)配合使用生成任务组handle。32图捕获阶段,与[torch_npu.npu.graph_task_group_end](torch_npu-npu-graph_task_group_end.md)配合使用生成任务组handle。
33 33 
34## 调用示例34## 调用示例
35+ 
35```python36```python
36import torch37import torch
37import torch_npu38import torch_npu
@@ -98,4 +99,3 @@ with torch.no_grad():
98 99 
99 100 
100```101```
101- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-graph_task_group_end.md+3-3
@@ -1,4 +1,5 @@
1# torch_npu.npu.graph_task_group_end1# torch_npu.npu.graph_task_group_end
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -8,14 +9,13 @@
8|<term>Atlas A2 训练系列产品</term> | √ |9|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas A2 推理系列产品</term> | √ |10|<term>Atlas A2 推理系列产品</term> | √ |
10 11 
11- 
12## 功能说明12## 功能说明
13 13 
14NPUGraph场景下,用于标记任务组结束位置。14NPUGraph场景下,用于标记任务组结束位置。
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.npu.graph_task_group_end(stream) -> handle19torch_npu.npu.graph_task_group_end(stream) -> handle
20```20```
21 21 
@@ -33,6 +33,7 @@ torch_npu.npu.graph_task_group_end(stream) -> handle
33- 入参stream需要同[torch_npu.npu.graph_task_group_begin](torch_npu-npu-graph_task_group_begin.md)保持一致。33- 入参stream需要同[torch_npu.npu.graph_task_group_begin](torch_npu-npu-graph_task_group_begin.md)保持一致。
34 34 
35## 调用示例35## 调用示例
36+ 
36```python37```python
37import torch38import torch
38import torch_npu39import torch_npu
@@ -99,4 +100,3 @@ with torch.no_grad():
99 100 
100 101 
101```102```
102- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-graph_task_update_begin.md+3-3
@@ -1,4 +1,5 @@
1# torch_npu.npu.graph_task_update_begin1# torch_npu.npu.graph_task_update_begin
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -8,14 +9,13 @@
8|<term>Atlas A2 训练系列产品</term> | √ |9|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas A2 推理系列产品</term> | √ |10|<term>Atlas A2 推理系列产品</term> | √ |
10 11 
11- 
12## 功能说明12## 功能说明
13 13 
14NPUGraph场景下,用于标记待更新任务的起始。14NPUGraph场景下,用于标记待更新任务的起始。
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.npu.graph_task_update_begin(stream, handle) -> None19torch_npu.npu.graph_task_update_begin(stream, handle) -> None
20```20```
21 21 
@@ -35,6 +35,7 @@ torch_npu.npu.graph_task_update_begin(stream, handle) -> None
35- 图更新阶段的流同图捕获阶段的流必须不同。35- 图更新阶段的流同图捕获阶段的流必须不同。
36 36 
37## 调用示例37## 调用示例
38+ 
38```python39```python
39import torch40import torch
40import torch_npu41import torch_npu
@@ -101,4 +102,3 @@ with torch.no_grad():
101 102 
102 103 
103```104```
104- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-graph_task_update_end.md+3-3
@@ -1,4 +1,5 @@
1# torch_npu.npu.graph_task_update_end1# torch_npu.npu.graph_task_update_end
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -8,14 +9,13 @@
8|<term>Atlas A2 训练系列产品</term> | √ |9|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas A2 推理系列产品</term> | √ |10|<term>Atlas A2 推理系列产品</term> | √ |
10 11 
11- 
12## 功能说明12## 功能说明
13 13 
14NPUGraph场景下,用于标记待更新任务的结束。14NPUGraph场景下,用于标记待更新任务的结束。
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.npu.graph_task_update_end(stream) -> None19torch_npu.npu.graph_task_update_end(stream) -> None
20```20```
21 21 
@@ -34,6 +34,7 @@ torch_npu.npu.graph_task_update_end(stream) -> None
34- 入参stream需要同[torch_npu.npu.graph_task_update_begin](torch_npu-npu-graph_task_update_begin.md)保持一致。34- 入参stream需要同[torch_npu.npu.graph_task_update_begin](torch_npu-npu-graph_task_update_begin.md)保持一致。
35 35 
36## 调用示例36## 调用示例
37+ 
37```python38```python
38import torch39import torch
39import torch_npu40import torch_npu
@@ -100,4 +101,3 @@ with torch.no_grad():
100 101 
101 102 
102```103```
103- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-host_empty_cache.md+5-6
@@ -11,35 +11,34 @@
11|<term>Atlas 推理系列产品</term> | √ |11|<term>Atlas 推理系列产品</term> | √ |
12|<term>Atlas 训练系列产品</term> | √ |12|<term>Atlas 训练系列产品</term> | √ |
13 13 
14- 
15## 功能说明14## 功能说明
16 15 
17释放当前由缓存持有的所有未占用的host物理内存。16释放当前由缓存持有的所有未占用的host物理内存。
18 17 
19- 
20## 定义文件18## 定义文件
19+ 
21torch_npu/npu/memory.py20torch_npu/npu/memory.py
22 21 
23## 函数原型22## 函数原型
24 23 
25-```24+```python
26torch_npu.npu.host_empty_cache()25torch_npu.npu.host_empty_cache()
27```26```
28 27 
29## 参数说明28## 参数说明
29+ 
3030
31 31 
32## 返回值说明32## 返回值说明
33+ 
3334
34 35 
35## 约束说明36## 约束说明
36 37 
3738
38 39 
39- 
40## 调用示例40## 调用示例
41 41 
42- 
43```python42```python
44>>> import torch43>>> import torch
45>>> import torch_npu44>>> import torch_npu
@@ -47,4 +46,4 @@ torch_npu.npu.host_empty_cache()
47>>> del x46>>> del x
48>>> torch_npu.npu.host_empty_cache()47>>> torch_npu.npu.host_empty_cache()
49>>> print(torch_npu.npu.host_memory_stats())48>>> print(torch_npu.npu.host_memory_stats())
50-```49+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-matmul-allow_hf32.md+3-5
@@ -1,4 +1,5 @@
1# torch_npu.npu.matmul.allow_hf321# torch_npu.npu.matmul.allow_hf32
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -8,7 +9,6 @@
8|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
9|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
10 11 
11- 
12## 功能说明12## 功能说明
13 13 
14matmul类算子开启对hf32类型的支持能力。14matmul类算子开启对hf32类型的支持能力。
@@ -17,18 +17,17 @@ matmul类算子开启对hf32类型的支持能力。
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.matmul.allow_hf32 = bool21torch_npu.npu.matmul.allow_hf32 = bool
22```22```
23 23 
24- 
25## 参数说明24## 参数说明
26 25 
27输入bool值,默认值False。26输入bool值,默认值False。
28 27 
29## 返回值说明28## 返回值说明
30-`bool`
31 29 
30+`bool`
32 31 
33## 调用示例32## 调用示例
34 33 
@@ -44,4 +43,3 @@ True
44>>>torch_npu.npu.matmul.allow_hf3243>>>torch_npu.npu.matmul.allow_hf32
45False44False
46```45```
47- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-matmul-cubeMathType.md+0-2
@@ -9,7 +9,6 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15设置或查询matmul类算子的cube计算类型。14设置或查询matmul类算子的cube计算类型。
@@ -36,7 +35,6 @@ torch_npu.npu.matmul.cube_math_type = CubeMathType
36| CubeMathType.USE_HF32 | 3 | 使用HF32计算模式 |35| CubeMathType.USE_HF32 | 3 | 使用HF32计算模式 |
37| CubeMathType.FORCE_GRP_ACC_FOR_FP32 | 4 | 强制使用FP32分组精度 |36| CubeMathType.FORCE_GRP_ACC_FOR_FP32 | 4 | 强制使用FP32分组精度 |
38 37 
39- 
40## 返回值说明38## 返回值说明
41 39 
42返回`CubeMathType`枚举类型。40返回`CubeMathType`枚举类型。
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-mstx-mark.md+2-4
@@ -17,13 +17,12 @@
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.mstx.mark(message: str='None', stream=None, domain: str='default') -> none:21torch_npu.npu.mstx.mark(message: str='None', stream=None, domain: str='default') -> none:
22```22```
23 23 
24## 参数说明24## 参数说明
25 25 
26- 
27- **message** (`str`):可选参数,打点携带信息字符串指针,默认为None。传入的message字符串长度要求:26- **message** (`str`):可选参数,打点携带信息字符串指针,默认为None。传入的message字符串长度要求:
28 - MSPTI场景:不能超过255字节。27 - MSPTI场景:不能超过255字节。
29 - 非MSPTI场景:不能超过156字节。28 - 非MSPTI场景:不能超过156字节。
@@ -40,7 +39,6 @@ torch_npu.npu.mstx.mark(message: str='None', stream=None, domain: str='default')
40 39 
41以下是关键步骤的代码示例,不可直接拷贝编译运行,仅供参考。40以下是关键步骤的代码示例,不可直接拷贝编译运行,仅供参考。
42 41 
43- 
44```python42```python
45import torch43import torch
46import torch_npu44import torch_npu
@@ -59,4 +57,4 @@ with torch_npu.profiler.profile(
59 for step in range(steps):57 for step in range(steps):
60 train_one_step() # 用户代码,包含调用mstx接口58 train_one_step() # 用户代码,包含调用mstx接口
61 prof.step()59 prof.step()
62-```60+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-mstx-mstx_range.md+3-3
@@ -1,4 +1,4 @@
1-# torch_npu.npu.mstx.mstx_range1+# torch_npu.npu.mstx.mstx_range
2 2 
3## 产品支持情况3## 产品支持情况
4 4 
@@ -17,7 +17,7 @@ range装饰器,用来采集被装饰函数的range执行耗时。
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.mstx.mstx_range(msg: str='None', stream=None, domain: str='default')21torch_npu.npu.mstx.mstx_range(msg: str='None', stream=None, domain: str='default')
22```22```
23 23 
@@ -63,4 +63,4 @@ with torch_npu.profiler.profile(
63 prof.step()63 prof.step()
64 64 
65 65 
66-```66+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-mstx-range_end.md+2-2
@@ -17,7 +17,7 @@
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.mstx.range_end(range_id: int, domain: str='default') -> int:21torch_npu.npu.mstx.range_end(range_id: int, domain: str='default') -> int:
22```22```
23 23 
@@ -38,4 +38,4 @@ torch_npu.npu.mstx.range_end(range_id: int, domain: str='default') -> int:
38id = torch_npu.npu.mstx.range_start("dataloader", None) # 第二个入参设置None或者不设置,只记录Host侧range耗时38id = torch_npu.npu.mstx.range_start("dataloader", None) # 第二个入参设置None或者不设置,只记录Host侧range耗时
39dataloader()39dataloader()
40torch_npu.npu.mstx.range_end(id)40torch_npu.npu.mstx.range_end(id)
41-```41+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-mstx-range_start.md+2-2
@@ -17,7 +17,7 @@
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu.mstx.range_start(message: str='None', stream=None, domain: str='default') -> int:21torch_npu.npu.mstx.range_start(message: str='None', stream=None, domain: str='default') -> int:
22```22```
23 23 
@@ -43,4 +43,4 @@ range_id:用于标识该range;如果接口执行失败,返回0。
43id = torch_npu.npu.mstx.range_start("dataloader", None) # 第二个入参设置None或者不设置,只记录Host侧range耗时43id = torch_npu.npu.mstx.range_start("dataloader", None) # 第二个入参设置None或者不设置,只记录Host侧range耗时
44dataloader()44dataloader()
45torch_npu.npu.mstx.range_end(id)45torch_npu.npu.mstx.range_end(id)
46-```46+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-mstx.md+3-4
@@ -15,11 +15,11 @@
15 15 
16打点接口。16打点接口。
17 17 
18-用于为[torch_npu.profiler._ExperimentalConfig](torch_npu-profiler-_ExperimentalConfig.md)的mstx提供打点接口调用。18+用于为[torch_npu.profiler._ExperimentalConfig](../torch_npu-profiler/torch_npu-profiler-_ExperimentalConfig.md)的mstx提供打点接口调用。
19 19 
20## 函数原型20## 函数原型
21 21 
22-```22+```python
23torch_npu.npu.mstx()23torch_npu.npu.mstx()
24```24```
25 25 
@@ -35,9 +35,8 @@ torch_npu.npu.mstx()
35 35 
36以下是关键步骤的代码示例,不可直接拷贝编译运行,仅供参考。36以下是关键步骤的代码示例,不可直接拷贝编译运行,仅供参考。
37 37 
38- 
39```python38```python
40import torch39import torch
41import torch_npu40import torch_npu
42mstx_object = torch_npu.npu.mstx()41mstx_object = torch_npu.npu.mstx()
43-```42+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-set_deterministic_level.md+3-2
@@ -16,6 +16,7 @@
16该接口用于控制CANN侧强一致性功能。具体为重新配置CANN侧参数,参数详细说明可见:[aclSysParamOpt](https://www.hiascend.com/document/detail/zh/CANNCommunityEdition/850/API/appdevgapi/aclcppdevg_03_1393.html)。16该接口用于控制CANN侧强一致性功能。具体为重新配置CANN侧参数,参数详细说明可见:[aclSysParamOpt](https://www.hiascend.com/document/detail/zh/CANNCommunityEdition/850/API/appdevgapi/aclcppdevg_03_1393.html)。
17 17 
18实际level对应配置如下表所示:18实际level对应配置如下表所示:
19+ 
19| torch_npu.npu.set_deterministic_level配置 | 调用aclrtSetSysParamOpt配置 | 实际功能 | 与原生接口torch.use_deterministic_algorithms的对应关系 |20| torch_npu.npu.set_deterministic_level配置 | 调用aclrtSetSysParamOpt配置 | 实际功能 | 与原生接口torch.use_deterministic_algorithms的对应关系 |
20| :---------------------------------------- | :-------------------------------------------------------------------------------------------- | :--------- | :------------------------|21| :---------------------------------------- | :-------------------------------------------------------------------------------------------- | :--------- | :------------------------|
21| 0 | aclrtSetSysParamOpt(ACL_OPT_DETERMINISTIC, 0)<br>aclrtSetSysParamOpt(ACL_OPT_STRONG_CONSISTENCY, 0) | 关闭确定性 | torch.use_deterministic_algorithms(False) |22| 0 | aclrtSetSysParamOpt(ACL_OPT_DETERMINISTIC, 0)<br>aclrtSetSysParamOpt(ACL_OPT_STRONG_CONSISTENCY, 0) | 关闭确定性 | torch.use_deterministic_algorithms(False) |
@@ -26,7 +27,7 @@
26 27 
27## 函数原型28## 函数原型
28 29 
29-```30+```python
30torch_npu.npu.set_deterministic_level(level)31torch_npu.npu.set_deterministic_level(level)
31```32```
32 33 
@@ -56,4 +57,4 @@ torch_npu.npu.set_deterministic_level(level)
56import torch57import torch
57import torch_npu58import torch_npu
58torch_npu.npu.set_deterministic_level(2)59torch_npu.npu.set_deterministic_level(2)
59-```60+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-set_op_timeout_ms.md+2-2
@@ -16,7 +16,7 @@
16 16 
17## 函数原型17## 函数原型
18 18 
19-```19+```python
20torch_npu.npu.set_op_timeout_ms(timeout)20torch_npu.npu.set_op_timeout_ms(timeout)
21```21```
22 22 
@@ -39,4 +39,4 @@ import torch
39import torch_npu39import torch_npu
40 40 
41torch_npu.npu.set_op_timeout_ms(1000)41torch_npu.npu.set_op_timeout_ms(1000)
42-```42+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu-use_compatible_impl.md+2-2
@@ -18,7 +18,7 @@
18 18 
19## 函数原型19## 函数原型
20 20 
21-```21+```python
22torch_npu.npu.use_compatible_impl(is_enable)22torch_npu.npu.use_compatible_impl(is_enable)
23```23```
24 24 
@@ -48,4 +48,4 @@ shape = [100, 400]
48mode = "none"48mode = "none"
49input = torch.rand(shape, dtype=torch.float16).npu()49input = torch.rand(shape, dtype=torch.float16).npu()
50output = torch.nn.functional.gelu(input, approximate=mode)50output = torch.nn.functional.gelu(input, approximate=mode)
51-```51+```
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu_list.md+0-1
@@ -613,4 +613,3 @@
613</tr>613</tr>
614</tbody>614</tbody>
615</table>615</table>
616- 
Mdocs/zh/custom_APIs/torch_npu-npu/torch_npu-npu_npugraph_handlers.md+9-7
@@ -13,12 +13,13 @@
13本接口用于注册自定义算子处理器,使自定义算子支持NPU Graph的动态Shape更新与重放功能。在NPU Graph模式下,用户调用 `g.update()` 传入新的参数时,Ascend Extension for PyTorch框架通过注册的处理器将数据映射到算子输入位置。13本接口用于注册自定义算子处理器,使自定义算子支持NPU Graph的动态Shape更新与重放功能。在NPU Graph模式下,用户调用 `g.update()` 传入新的参数时,Ascend Extension for PyTorch框架通过注册的处理器将数据映射到算子输入位置。
14 14 
15核心机制:15核心机制:
16+ 
161. Capture预处理:定义算子捕获时的输入数据预处理逻辑。171. Capture预处理:定义算子捕获时的输入数据预处理逻辑。
172. Update动态更新:在Graph Replay(回放)阶段,无需重新Capture图结构,即可动态修改算子输入参数(如序列长度、Batch Size等)的机制。具体流程如下:182. Update动态更新:在Graph Replay(回放)阶段,无需重新Capture图结构,即可动态修改算子输入参数(如序列长度、Batch Size等)的机制。具体流程如下:
18- 1. 用户在Replay前调用`g.update(cpu_update_input=[...])`传入新参数。19+ 1. 用户在Replay前调用`g.update(cpu_update_input=[...])`传入新参数。
19- 2. 框架遍历Graph中的算子,查找注册的`NpuGraphOpHandler`。20+ 2. 框架遍历Graph中的算子,查找注册的`NpuGraphOpHandler`。
20- 3. 框架调用Handler的`update_args`方法,传入`dispatch_record` 和`update_input`。21+ 3. 框架调用Handler的`update_args`方法,传入`dispatch_record` 和`update_input`。
21- 4. 在`update_args`中,用户直接修改`dispatch_record.args`指定索引的值。22+ 4. 在`update_args`中,用户直接修改`dispatch_record.args`指定索引的值。
22 23 
23## 定义文件24## 定义文件
24 25 
@@ -51,15 +52,18 @@ def register_npu_graph_handler(op_names: str | list[str]): ...
51### NpuGraphOpHandler基类52### NpuGraphOpHandler基类
52 53 
53#### 1. prepare_capture54#### 1. prepare_capture
55+ 
54- **func**`Callable`):原始算子函数对象。56- **func**`Callable`):原始算子函数对象。
55- **args**`Tuple`):位置参数元组。57- **args**`Tuple`):位置参数元组。
56- **kwargs**`Dict`):关键字参数字典。58- **kwargs**`Dict`):关键字参数字典。
57 59 
58#### 2. postprocess_result60#### 2. postprocess_result
61+ 
59- **result**`Any`):算子执行的原始结果。62- **result**`Any`):算子执行的原始结果。
60- **kwargs**`Dict`):关键字参数字典。63- **kwargs**`Dict`):关键字参数字典。
61 64 
62#### 3. update_args65#### 3. update_args
66+ 
63- **dispatch_record**(`DispatchRecord`):包含算子运行时信息的记录对象,可通过 `dispatch_record.args` 修改参数。67- **dispatch_record**(`DispatchRecord`):包含算子运行时信息的记录对象,可通过 `dispatch_record.args` 修改参数。
64- **update_input**(`Dict`):用户传入的更新参数字典。68- **update_input**(`Dict`):用户传入的更新参数字典。
65 69 
@@ -69,7 +73,7 @@ def register_npu_graph_handler(op_names: str | list[str]): ...
69- **value**`Any`):关键字参数的值。73- **value**`Any`):关键字参数的值。
70- **tensor_param_names**`List[str]`):张量参数名称列表。74- **tensor_param_names**`List[str]`):张量参数名称列表。
71 75 
72-> [!CAUTION]</br>76+> [!CAUTION]<br>
73> 所有方法必须声明为 `classmethod`,禁止使用 `self` 存储状态(无状态设计)。77> 所有方法必须声明为 `classmethod`,禁止使用 `self` 存储状态(无状态设计)。
74 78 
75### register_npu_graph_handler装饰器79### register_npu_graph_handler装饰器
@@ -102,7 +106,6 @@ def register_npu_graph_handler(op_names: str | list[str]): ...
102 106 
103## 调用示例107## 调用示例
104 108 
105- 
106本示例展示如何自定义一个Handler,同时实现输出预分配(`prepare_capture`)、返回值格式调整(`postprocess_result`)、动态参数更新(`update_args`)以及kwargs自定义存储(`record_wrap_kwarg`)。109本示例展示如何自定义一个Handler,同时实现输出预分配(`prepare_capture`)、返回值格式调整(`postprocess_result`)、动态参数更新(`update_args`)以及kwargs自定义存储(`record_wrap_kwarg`)。
107 110 
108```python111```python
@@ -161,4 +164,3 @@ with torch.npu.graph(g, auto_dispatch_capture=True):
161g.update(cpu_update_input=[{"seq_len": new_seq_len}])164g.update(cpu_update_input=[{"seq_len": new_seq_len}])
162g.replay()165g.replay()
163```166```
164- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)NPU-Tensor.md+0-1
@@ -82,4 +82,3 @@ Torch_npu提供NPU tensor相关的部分接口使用与Cuda类似。
82</tr>82</tr>
83</tbody>83</tbody>
84</table>84</table>
85- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-aclnn-version.md+1-3
@@ -9,14 +9,13 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15查询aclnn算子版本信息。aclnn算子详情可参考《CANN AOL算子加速库接口》中的“<a href="https://www.hiascend.com/document/detail/zh/canncommercial/850/API/aolapi/operatorlist_00001.html">接口简介</a>”章节。14查询aclnn算子版本信息。aclnn算子详情可参考《CANN AOL算子加速库接口》中的“<a href="https://www.hiascend.com/document/detail/zh/canncommercial/850/API/aolapi/operatorlist_00001.html">接口简介</a>”章节。
16 15 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.npu.aclnn.version(): -> None19torch_npu.npu.aclnn.version(): -> None
21```20```
22 21 
@@ -31,4 +30,3 @@ torch_npu.npu.aclnn.version(): -> None
31>>> import torch_npu30>>> import torch_npu
32>>> res = torch_npu.npu.aclnn.version()31>>> res = torch_npu.npu.aclnn.version()
33```32```
34- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-check_uce_in_memory.md+1-5
@@ -10,7 +10,6 @@
10|<term>Atlas A3 训练系列产品</term> | √ |10|<term>Atlas A3 训练系列产品</term> | √ |
11|<term>Atlas A2 训练系列产品</term> | √ |11|<term>Atlas A2 训练系列产品</term> | √ |
12 12 
13- 
14## 功能说明13## 功能说明
15 14 
16提供故障内存地址类型检测接口,供MindCluster进行故障恢复策略的决策。其功能是在出现UCE片上内存故障时,判断故障内存地址类型。15提供故障内存地址类型检测接口,供MindCluster进行故障恢复策略的决策。其功能是在出现UCE片上内存故障时,判断故障内存地址类型。
@@ -20,7 +19,7 @@
20 19 
21## 函数原型20## 函数原型
22 21 
23-```22+```python
24torch_npu.npu.check_uce_in_memory(device_id:int)23torch_npu.npu.check_uce_in_memory(device_id:int)
25```24```
26 25 
@@ -28,7 +27,6 @@ torch_npu.npu.check_uce_in_memory(device_id:int)
28 27 
29**device_id** (`int`):需要处理的device id。28**device_id** (`int`):需要处理的device id。
30 29 
31- 
32## 返回值说明30## 返回值说明
33 31 
34- 0:无UCE故障地址。32- 0:无UCE故障地址。
@@ -36,7 +34,6 @@ torch_npu.npu.check_uce_in_memory(device_id:int)
36- 2:UCE故障地址为Ascend Extension for PyTorch使用的临时内存地址。34- 2:UCE故障地址为Ascend Extension for PyTorch使用的临时内存地址。
37- 3:UCE故障地址为Ascend Extension for PyTorch使用的常驻内存地址。35- 3:UCE故障地址为Ascend Extension for PyTorch使用的常驻内存地址。
38 36 
39- 
40## 约束说明37## 约束说明
41 38 
42要确保是一个有效的device。39要确保是一个有效的device。
@@ -48,4 +45,3 @@ torch_npu.npu.check_uce_in_memory(device_id:int)
48>>> torch.npu.set_device(0)45>>> torch.npu.set_device(0)
49>>> torch_npu.npu.check_uce_in_memory(0)46>>> torch_npu.npu.check_uce_in_memory(0)
50 ```47 ```
51- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-clear_npu_overflow_flag.md+3-1
@@ -1,4 +1,5 @@
1# (beta)torch\_npu.npu.clear\_npu\_overflow\_flag1# (beta)torch\_npu.npu.clear\_npu\_overflow\_flag
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -10,6 +11,7 @@
10对NPU溢出检测进行清零。11对NPU溢出检测进行清零。
11 12 
12## 函数原型13## 函数原型
13-```14+ 
15+```python
14torch_npu.npu.clear_npu_overflow_flag()16torch_npu.npu.clear_npu_overflow_flag()
15```17```
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-config-allow_internal_format.md+3-2
@@ -1,4 +1,5 @@
1# (beta)torch_npu.npu.config.allow_internal_format1# (beta)torch_npu.npu.config.allow_internal_format
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -14,13 +15,14 @@
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.npu.config.allow_internal_format = bool19torch_npu.npu.config.allow_internal_format = bool
19```20```
20 21 
21## 参数说明22## 参数说明
22 23 
23输入`bool`值。24输入`bool`值。
25+ 
24- <term>Atlas A2 训练系列产品</term>/<term>Atlas A3 训练系列产品</term>默认值为`False`26- <term>Atlas A2 训练系列产品</term>/<term>Atlas A3 训练系列产品</term>默认值为`False`
25- <term>Atlas 推理系列产品</term>/<term>Atlas 训练系列产品</term>默认值为`True`27- <term>Atlas 推理系列产品</term>/<term>Atlas 训练系列产品</term>默认值为`True`
26 28 
@@ -35,4 +37,3 @@ torch_npu.npu.config.allow_internal_format = bool
35>>> import torch_npu37>>> import torch_npu
36>>> torch_npu.npu.config.allow_internal_format = False38>>> torch_npu.npu.config.allow_internal_format = False
37```39```
38- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-finalize_dump.md+1-4
@@ -9,15 +9,12 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13- 
14## 功能说明12## 功能说明
15 13 
16结束dump。配置前需要确保环境变量`NPU_DUMP_ENABLE=1`已设置。14结束dump。配置前需要确保环境变量`NPU_DUMP_ENABLE=1`已设置。
17 15 
18## 函数原型16## 函数原型
19 17 
20-```18+```python
21torch_npu.npu.finalize_dump()19torch_npu.npu.finalize_dump()
22```20```
23- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-get_amp_supported_dtype.md+1-4
@@ -9,15 +9,13 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13- 
14## 功能说明12## 功能说明
15 13 
16获取npu设备支持的数据类型,可能设备支持不止一种数据类型。14获取npu设备支持的数据类型,可能设备支持不止一种数据类型。
17 15 
18## 函数原型16## 函数原型
19 17 
20-```18+```python
21torch_npu.npu.get_amp_supported_dtype()19torch_npu.npu.get_amp_supported_dtype()
22```20```
23 21 
@@ -35,4 +33,3 @@ supported_dtypes = torch_npu.npu.get_amp_supported_dtype()
35print(f"NPU支持的AMP数据类型:{supported_dtypes}")33print(f"NPU支持的AMP数据类型:{supported_dtypes}")
36 34 
37```35```
38- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-get_autocast_dtype.md+2-5
@@ -1,6 +1,5 @@
1# (beta)torch_npu.npu.get_autocast_dtype1# (beta)torch_npu.npu.get_autocast_dtype
2 2 
3- 
4## 产品支持情况3## 产品支持情况
5 4 
6| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -10,17 +9,16 @@
10|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
11|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
12 11 
13- 
14## 功能说明12## 功能说明
15 13 
16在amp场景获取设备支持的数据类型,该`dtype`由torch_npu.npu.set_autocast_dtype设置或者默认数据类型`float16`。14在amp场景获取设备支持的数据类型,该`dtype`由torch_npu.npu.set_autocast_dtype设置或者默认数据类型`float16`。
17 15 
18- 
19## 函数原型16## 函数原型
20 17 
21-```18+```python
22torch_npu.npu.get_autocast_dtype()19torch_npu.npu.get_autocast_dtype()
23```20```
21+ 
24## 返回值说明22## 返回值说明
25 23 
26`torch.dtype`24`torch.dtype`
@@ -34,4 +32,3 @@ import torch_npu
34current_dtype = torch_npu.npu.get_autocast_dtype()32current_dtype = torch_npu.npu.get_autocast_dtype()
35 33 
36```34```
37- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-get_mm_bmm_format_nd.md+2-3
@@ -1,4 +1,5 @@
1# (beta)torch_npu.npu.get_mm_bmm_format_nd1# (beta)torch_npu.npu.get_mm_bmm_format_nd
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -14,7 +15,7 @@
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.npu.get_mm_bmm_format_nd()19torch_npu.npu.get_mm_bmm_format_nd()
19```20```
20 21 
@@ -22,7 +23,6 @@ torch_npu.npu.get_mm_bmm_format_nd()
22 23 
23`bool`24`bool`
24 25 
25- 
26## 调用示例26## 调用示例
27 27 
28```python28```python
@@ -31,4 +31,3 @@ torch_npu.npu.get_mm_bmm_format_nd()
31>>> torch_npu.npu.get_mm_bmm_format_nd()31>>> torch_npu.npu.get_mm_bmm_format_nd()
32True32True
33```33```
34- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-get_npu_overflow_flag.md+2-2
@@ -1,4 +1,5 @@
1# (beta)torch\_npu.npu.get\_npu\_overflow\_flag1# (beta)torch\_npu.npu.get\_npu\_overflow\_flag
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -11,7 +12,7 @@
11 12 
12## 函数原型13## 函数原型
13 14 
14-```15+```python
15torch_npu.npu.get_npu_overflow_flag()16torch_npu.npu.get_npu_overflow_flag()
16```17```
17 18 
@@ -24,4 +25,3 @@ a = torch.Tensor([65535]).npu().half()
24a = a + a25a = a + a
25ret = torch_npu.npu.get_npu_overflow_flag()26ret = torch_npu.npu.get_npu_overflow_flag()
26```27```
27- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-init_dump.md+1-3
@@ -9,14 +9,12 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15初始化dump配置。配置前需要确保环境变量`NPU_DUMP_ENABLE=1`已设置,以及已通过`torch_npu.npu.set_dump_config(path="/tmp/dump", mode="all")`配置dump。14初始化dump配置。配置前需要确保环境变量`NPU_DUMP_ENABLE=1`已设置,以及已通过`torch_npu.npu.set_dump_config(path="/tmp/dump", mode="all")`配置dump。
16 15 
17- 
18## 函数原型16## 函数原型
19 17 
20-```18+```python
21torch_npu.npu.init_dump()19torch_npu.npu.init_dump()
22```20```
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-is_autocast_enabled.md+2-5
@@ -9,22 +9,20 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15确认autocast是否可用。14确认autocast是否可用。
16 15 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.npu.is_autocast_enabled()19torch_npu.npu.is_autocast_enabled()
21```20```
21+ 
22## 返回值说明22## 返回值说明
23 23 
24`bool`24`bool`
25 25 
26- 
27- 
28## 调用示例26## 调用示例
29 27 
30``` python28``` python
@@ -32,4 +30,3 @@ import torch
32import torch_npu30import torch_npu
33torch_npu.npu.is_autocast_enabled()31torch_npu.npu.is_autocast_enabled()
34```32```
35- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-is_jit_compile_false.md+2-3
@@ -9,16 +9,16 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15确认JIT编译模式是否被禁用,如果被禁用,返回True,否则返回False。14确认JIT编译模式是否被禁用,如果被禁用,返回True,否则返回False。
16 15 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.npu.is_jit_compile_false()19torch_npu.npu.is_jit_compile_false()
21```20```
21+ 
22## 返回值说明22## 返回值说明
23 23 
24bool型。24bool型。
@@ -32,4 +32,3 @@ torch_npu.npu.set_compile_mode(jit_compile=False)
32torch_npu.npu.is_jit_compile_false()32torch_npu.npu.is_jit_compile_false()
33True33True
34```34```
35- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-obfuscation_calculate.md+7-5
@@ -1,5 +1,4 @@
1 1 
2- 
3# (beta)torch_npu.npu.obfuscation_calculate2# (beta)torch_npu.npu.obfuscation_calculate
4 3 
5## 产品支持情况4## 产品支持情况
@@ -13,10 +12,12 @@
13 12 
14该接口用于将张量`x`和配置参数(如`param`)发送至PMCC(Privacy and Model Confidential Computing)混淆引擎。引擎的CA(普通OS中的Client Application)模块调用TA(TEE OS中的Trusted Application)模块,进行张量混淆处理,最终返回混淆结果。13该接口用于将张量`x`和配置参数(如`param`)发送至PMCC(Privacy and Model Confidential Computing)混淆引擎。引擎的CA(普通OS中的Client Application)模块调用TA(TEE OS中的Trusted Application)模块,进行张量混淆处理,最终返回混淆结果。
15该接口针对[PMCC](https://www-file.huawei.com/admin/asset/v1/pro/view/6812dab6dd4e4640b11619e401db1c47.pdf)业务,如下两种结果均符合预期:14该接口针对[PMCC](https://www-file.huawei.com/admin/asset/v1/pro/view/6812dab6dd4e4640b11619e401db1c47.pdf)业务,如下两种结果均符合预期:
15+ 
16* 如果部署PMCC特性,该接口返回响应结果。16* 如果部署PMCC特性,该接口返回响应结果。
17* 未部署PMCC特性时,执行用例会返回错误码507018。17* 未部署PMCC特性时,执行用例会返回错误码507018。
18 18 
19PMCC特性的部署流程如下:19PMCC特性的部署流程如下:
20+ 
201. 环境中存在NPU驱动和固件。211. 环境中存在NPU驱动和固件。
212. 安装AI混淆SDK,执行一键式部署脚本,该脚本会自动完成以下任务:222. 安装AI混淆SDK,执行一键式部署脚本,该脚本会自动完成以下任务:
22 * 配置kmsAgent。23 * 配置kmsAgent。
@@ -29,20 +30,21 @@ PMCC特性的详细部署流程请参考对应的部署指导手册。
29 30 
30## 函数原型31## 函数原型
31 32 
32-```33+```python
33torch_npu.npu.obfuscation_calculate(fd, x, param, obf_coefficient) -> Tensor34torch_npu.npu.obfuscation_calculate(fd, x, param, obf_coefficient) -> Tensor
34```35```
35 36 
36## 参数说明37## 参数说明
37 38 
38-- **fd**(`Tensor`):必选参数,socket连接符,数据类型为`int32`,填写调用[obfuscation_initialize](./torch_npu-npu-obfuscation_initialize.md)接口的返回值。39+- **fd**(`Tensor`):必选参数,socket连接符,数据类型为`int32`,填写调用[obfuscation_initialize]((beta)torch_npu-npu-obfuscation_initialize.md)接口的返回值。
39-- **x**(`Tensor`):必选参数,待混淆处理的`Tensor`输入,对`Tensor`维度不作限制,shape为( , *, ... , hiddenSize),即最后一维的size是[obfuscation_initialize](./torch_npu-npu-obfuscation_initialize.md)的入参`hiddenSize`。数据格式支持ND。40+- **x**(`Tensor`):必选参数,待混淆处理的`Tensor`输入,对`Tensor`维度不作限制,shape为( , *, ... , hiddenSize),即最后一维的size是[obfuscation_initialize]((beta)torch_npu-npu-obfuscation_initialize.md)的入参`hiddenSize`。数据格式支持ND。
40 * <term>Atlas 推理系列产品</term>: 数据类型支持`float16``float32``int8`41 * <term>Atlas 推理系列产品</term>: 数据类型支持`float16``float32``int8`
41 * <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>: 数据类型支持`float16``float32``bfloat16``int8`42 * <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>: 数据类型支持`float16``float32``bfloat16``int8`
42- **param**`Tensor`):必选参数,张量`x`的最后一维的维度,数据类型为`int32`43- **param**`Tensor`):必选参数,张量`x`的最后一维的维度,数据类型为`int32`
43- **obf_coefficient**(`float`):可选参数,混淆系数,支持输入范围为(0.0,1.0],默认值1.0。44- **obf_coefficient**(`float`):可选参数,混淆系数,支持输入范围为(0.0,1.0],默认值1.0。
44 45 
45## 返回值说明46## 返回值说明
47+ 
46`Tensor`48`Tensor`
47 49 
48代表`obfuscation_calculate`的计算结果,输出数据类型及shape与`x`相同。50代表`obfuscation_calculate`的计算结果,输出数据类型及shape与`x`相同。
@@ -67,4 +69,4 @@ obf_cft = 1.0
67fd = torch_npu.npu.obfuscation_initialize(hidden_size, tp_rank, cmd, data_type=data_type, thread_num= thread_num, obf_coefficient=obf_cft)69fd = torch_npu.npu.obfuscation_initialize(hidden_size, tp_rank, cmd, data_type=data_type, thread_num= thread_num, obf_coefficient=obf_cft)
68param = torch.tensor([3584], device=device)70param = torch.tensor([3584], device=device)
69x_obf_out = torch_npu.npu.obfuscation_calculate(fd, hidden_states, param, obf_coefficient=obf_cft)71x_obf_out = torch_npu.npu.obfuscation_calculate(fd, hidden_states, param, obf_coefficient=obf_cft)
70-```72+```
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-obfuscation_finalize.md+6-4
@@ -1,5 +1,4 @@
1 1 
2- 
3# (beta)torch_npu.npu.obfuscation_finalize2# (beta)torch_npu.npu.obfuscation_finalize
4 3 
5## 产品支持情况4## 产品支持情况
@@ -13,10 +12,12 @@
13 12 
14该接口用于完成PMCC(Privacy and Model Confidential Computing)模型混淆引擎的资源释放,即与PMCC混淆引擎CA(普通OS中的Client Application)断开socket连接。13该接口用于完成PMCC(Privacy and Model Confidential Computing)模型混淆引擎的资源释放,即与PMCC混淆引擎CA(普通OS中的Client Application)断开socket连接。
15该接口针对[PMCC](https://www-file.huawei.com/admin/asset/v1/pro/view/6812dab6dd4e4640b11619e401db1c47.pdf)业务,如下两种结果均符合预期:14该接口针对[PMCC](https://www-file.huawei.com/admin/asset/v1/pro/view/6812dab6dd4e4640b11619e401db1c47.pdf)业务,如下两种结果均符合预期:
15+ 
16* 如果部署PMCC特性,该接口返回响应结果。16* 如果部署PMCC特性,该接口返回响应结果。
17* 未部署PMCC特性时,执行用例会返回错误码507018。17* 未部署PMCC特性时,执行用例会返回错误码507018。
18 18 
19PMCC特性的部署流程如下:19PMCC特性的部署流程如下:
20+ 
201. 环境中存在NPU驱动和固件。211. 环境中存在NPU驱动和固件。
212. 安装AI混淆SDK,执行一键式部署脚本,该脚本会自动完成以下任务:222. 安装AI混淆SDK,执行一键式部署脚本,该脚本会自动完成以下任务:
22 * 配置kmsAgent。23 * 配置kmsAgent。
@@ -29,15 +30,16 @@ PMCC特性的详细部署流程请参考对应的部署指导手册。
29 30 
30## 函数原型31## 函数原型
31 32 
32-```33+```python
33torch_npu.npu.obfuscation_finalize(fd_to_close) -> Tensor34torch_npu.npu.obfuscation_finalize(fd_to_close) -> Tensor
34```35```
35 36 
36## 参数说明37## 参数说明
37 38 
38-**fd_to_close**(`Tensor`):填写调用[obfuscation_initialize](./torch_npu-npu-obfuscation_initialize.md)接口的返回值,数据类型为`int32`。39+**fd_to_close**(`Tensor`):填写调用[obfuscation_initialize]((beta)torch_npu-npu-obfuscation_initialize.md)接口的返回值,数据类型为`int32`。
39 40 
40## 返回值说明41## 返回值说明
42+ 
41`Tensor`43`Tensor`
42 44 
43代表关闭socket连接符内存数据,1D,shape为(1),数据类型为`int32`45代表关闭socket连接符内存数据,1D,shape为(1),数据类型为`int32`
@@ -61,4 +63,4 @@ hidden_states = torch.randn((1024,3584), dtype=torch.bfloat16, device=device)
61obf_cft = 1.063obf_cft = 1.0
62fd = torch_npu.npu.obfuscation_initialize(hidden_size, tp_rank, cmd, data_type=data_type, thread_num= thread_num, obf_coefficient=obf_cft)64fd = torch_npu.npu.obfuscation_initialize(hidden_size, tp_rank, cmd, data_type=data_type, thread_num= thread_num, obf_coefficient=obf_cft)
63torch_npu.npu.obfuscation_finalize(fd)65torch_npu.npu.obfuscation_finalize(fd)
64-```66+```
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-obfuscation_initialize.md+5-3
@@ -1,5 +1,4 @@
1 1 
2- 
3# (beta)torch_npu.npu.obfuscation_initialize2# (beta)torch_npu.npu.obfuscation_initialize
4 3 
5## 产品支持情况4## 产品支持情况
@@ -13,10 +12,12 @@
13 12 
14该接口用于完成PMCC(Privacy and Model Confidential Computing)模型混淆引擎的资源初始化,即与PMCC混淆引擎CA(普通OS中的Client Application)建立socket连接、对CA、TA(TEE OS中的Trusted Application)进行初始化,并返回socket连接符。13该接口用于完成PMCC(Privacy and Model Confidential Computing)模型混淆引擎的资源初始化,即与PMCC混淆引擎CA(普通OS中的Client Application)建立socket连接、对CA、TA(TEE OS中的Trusted Application)进行初始化,并返回socket连接符。
15该接口针对[PMCC](https://www-file.huawei.com/admin/asset/v1/pro/view/6812dab6dd4e4640b11619e401db1c47.pdf)业务,如下两种结果均符合预期:14该接口针对[PMCC](https://www-file.huawei.com/admin/asset/v1/pro/view/6812dab6dd4e4640b11619e401db1c47.pdf)业务,如下两种结果均符合预期:
15+ 
16* 如果部署PMCC特性,该接口返回响应结果。16* 如果部署PMCC特性,该接口返回响应结果。
17* 未部署PMCC特性时,执行用例会返回错误码507018。17* 未部署PMCC特性时,执行用例会返回错误码507018。
18 18 
19PMCC特性的部署流程如下:19PMCC特性的部署流程如下:
20+ 
201. 环境中存在NPU驱动和固件。211. 环境中存在NPU驱动和固件。
212. 安装AI混淆SDK,执行一键式部署脚本,该脚本会自动完成以下任务:222. 安装AI混淆SDK,执行一键式部署脚本,该脚本会自动完成以下任务:
22 * 配置kmsAgent。23 * 配置kmsAgent。
@@ -29,7 +30,7 @@ PMCC特性的详细部署流程请参考对应的部署指导手册。
29 30 
30## 函数原型31## 函数原型
31 32 
32-```33+```python
33torch_npu.npu.obfuscation_initialize(hidden_size, tp_rank, cmd, data_type, model_obf_seed_id, data_obf_seed_id, thread_num, obf_coefficient) -> Tensor34torch_npu.npu.obfuscation_initialize(hidden_size, tp_rank, cmd, data_type, model_obf_seed_id, data_obf_seed_id, thread_num, obf_coefficient) -> Tensor
34```35```
35 36 
@@ -50,6 +51,7 @@ torch_npu.npu.obfuscation_initialize(hidden_size, tp_rank, cmd, data_type, model
50- **obf_coefficient**(`float`):可选参数,混淆系数,支持输入范围为0-1,默认值1.0。51- **obf_coefficient**(`float`):可选参数,混淆系数,支持输入范围为0-1,默认值1.0。
51 52 
52## 返回值说明53## 返回值说明
54+ 
53`Tensor`55`Tensor`
54 56 
55代表socket连接符,1D,shape为(1),数据类型为`int32`57代表socket连接符,1D,shape为(1),数据类型为`int32`
@@ -72,4 +74,4 @@ i = 0
72hidden_states = torch.randn((1024,3584), dtype=torch.bfloat16, device=device)74hidden_states = torch.randn((1024,3584), dtype=torch.bfloat16, device=device)
73obf_cft = 1.075obf_cft = 1.0
74fd = torch_npu.npu.obfuscation_initialize(hidden_size, tp_rank, cmd, data_type=data_type, thread_num= thread_num, obf_coefficient=obf_cft)76fd = torch_npu.npu.obfuscation_initialize(hidden_size, tp_rank, cmd, data_type=data_type, thread_num= thread_num, obf_coefficient=obf_cft)
75-```77+```
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-restart_device.md+1-3
@@ -9,7 +9,7 @@
9 9 
10## 函数原型10## 函数原型
11 11 
12-```12+```python
13torch_npu.npu.restart_device(device_id: int, rebuild_all_resource: bool = False) -> None13torch_npu.npu.restart_device(device_id: int, rebuild_all_resource: bool = False) -> None
14```14```
15 15 
@@ -22,7 +22,6 @@ torch_npu.npu.restart_device(device_id: int, rebuild_all_resource: bool = False)
22 22 
23要确保是一个有效的device,这个device可以是被stop过的也可以是没有被stop。23要确保是一个有效的device,这个device可以是被stop过的也可以是没有被stop。
24 24 
25- 
26## 调用示例25## 调用示例
27 26 
28 ```python27 ```python
@@ -32,4 +31,3 @@ torch_npu.npu.restart_device(device_id: int, rebuild_all_resource: bool = False)
32>>> torch_npu.npu.stop_device(0)31>>> torch_npu.npu.stop_device(0)
33>>> torch_npu.npu.restart_device(0)32>>> torch_npu.npu.restart_device(0)
34 ```33 ```
35- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-set_aoe.md+1-1
@@ -15,7 +15,7 @@ AOE调优使能。
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.npu.set_aoe(dump_path)19torch_npu.npu.set_aoe(dump_path)
20```20```
21 21 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-set_autocast_dtype.md+1-4
@@ -9,14 +9,13 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15设置设备在AMP场景支持的数据类型。14设置设备在AMP场景支持的数据类型。
16 15 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.npu.set_autocast_dtype(dtype)19torch_npu.npu.set_autocast_dtype(dtype)
21```20```
22 21 
@@ -24,7 +23,6 @@ torch_npu.npu.set_autocast_dtype(dtype)
24 23 
25 **dtype** :数据类型。24 **dtype** :数据类型。
26 25 
27- 
28## 调用示例26## 调用示例
29 27 
30```python28```python
@@ -32,4 +30,3 @@ torch_npu.npu.set_autocast_dtype(dtype)
32>>> import torch_npu30>>> import torch_npu
33>>> torch_npu.npu.set_autocast_dtype(torch.float16)31>>> torch_npu.npu.set_autocast_dtype(torch.float16)
34```32```
35- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-set_autocast_enabled.md+1-4
@@ -9,14 +9,13 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15是否在设备上使能AMP。14是否在设备上使能AMP。
16 15 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.npu.set_autocast_enabled(bool)19torch_npu.npu.set_autocast_enabled(bool)
21```20```
22 21 
@@ -24,7 +23,6 @@ torch_npu.npu.set_autocast_enabled(bool)
24 23 
25**bool** :入参为True时,在设备上使能AMP,否则,不使能AMP。24**bool** :入参为True时,在设备上使能AMP,否则,不使能AMP。
26 25 
27- 
28## 调用示例26## 调用示例
29 27 
30```python28```python
@@ -32,4 +30,3 @@ import torch
32import torch_npu30import torch_npu
33torch_npu.npu.set_autocast_enabled(True)31torch_npu.npu.set_autocast_enabled(True)
34```32```
35- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-set_compile_mode.md+3-4
@@ -9,28 +9,27 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13## 功能说明12## 功能说明
14 13 
15设置是否开启二进制。14设置是否开启二进制。
16 15 
17## 函数原型16## 函数原型
18 17 
19-```18+```python
20torch_npu.npu.set_compile_mode(jit_compile = bool)19torch_npu.npu.set_compile_mode(jit_compile = bool)
21```20```
21+ 
22## 参数说明22## 参数说明
23 23 
24**jit_compile**(`bool`):设置为True时表示非二进制模式,设置为False时表示二进制模式。24**jit_compile**(`bool`):设置为True时表示非二进制模式,设置为False时表示二进制模式。
25 25 
26> [!NOTE] 26> [!NOTE]
27+>
27>- Atlas 训练系列产品/Atlas 推理系列产品默认为jit_compile=True,即非二进制模式。28>- Atlas 训练系列产品/Atlas 推理系列产品默认为jit_compile=True,即非二进制模式。
28>- Atlas A2 训练系列产品/Atlas A3 训练系列产品默认为jit_compile=False,即二进制模式。29>- Atlas A2 训练系列产品/Atlas A3 训练系列产品默认为jit_compile=False,即二进制模式。
29 30 
30- 
31## 调用示例31## 调用示例
32 32 
33```python33```python
34>>> torch_npu.npu.set_compile_mode(jit_compile=False)34>>> torch_npu.npu.set_compile_mode(jit_compile=False)
35```35```
36- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-set_dump.md+1-4
@@ -15,16 +15,14 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.npu.set_dump(path_to_json)19torch_npu.npu.set_dump(path_to_json)
20```20```
21 21 
22- 
23## 参数说明22## 参数说明
24 23 
25 **path_to_json**:配置文件所在的路径,包含文件名,用户需根据实际情况配置。具体配置请参考《CANN 应用开发接口 (Python)》中“<a href="https://www.hiascend.com/document/detail/zh/canncommercial/850/API/appdevgapi/aclpythondevg_01_0155.html">函数:set_dump</a>”章节。24 **path_to_json**:配置文件所在的路径,包含文件名,用户需根据实际情况配置。具体配置请参考《CANN 应用开发接口 (Python)》中“<a href="https://www.hiascend.com/document/detail/zh/canncommercial/850/API/appdevgapi/aclpythondevg_01_0155.html">函数:set_dump</a>”章节。
26 25 
27- 
28## 调用示例26## 调用示例
29 27 
30```python28```python
@@ -32,4 +30,3 @@ torch_npu.npu.set_dump(path_to_json)
32>>> import torch_npu30>>> import torch_npu
33>>> torch_npu.npu.set_dump("/home/HwHiAiUser/dump.json")31>>> torch_npu.npu.set_dump("/home/HwHiAiUser/dump.json")
34```32```
35- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-set_mm_bmm_format_nd.md+2-1
@@ -1,4 +1,5 @@
1# (beta)torch_npu.npu.set_mm_bmm_format_nd1# (beta)torch_npu.npu.set_mm_bmm_format_nd
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -14,7 +15,7 @@
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.npu.set_mm_bmm_format_nd(bool)19torch_npu.npu.set_mm_bmm_format_nd(bool)
19```20```
20 21 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-set_option.md+0-1
@@ -1,4 +1,3 @@
1# (beta)torch_npu.npu.set_option1# (beta)torch_npu.npu.set_option
2 2 
3详细使用参见《PyTorch 训练模型迁移调优指南》中的“<a href="https://www.hiascend.com/document/detail/zh/Pytorch/730/ptmoddevg/trainingmigrguide/PT_LMTMOG_0076.html">设置算子编译选项</a>”章节。3详细使用参见《PyTorch 训练模型迁移调优指南》中的“<a href="https://www.hiascend.com/document/detail/zh/Pytorch/730/ptmoddevg/trainingmigrguide/PT_LMTMOG_0076.html">设置算子编译选项</a>”章节。
4- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-stop_device.md+2-2
@@ -9,7 +9,7 @@
9 9 
10## 函数原型10## 函数原型
11 11 
12-```12+```python
13torch_npu.npu.stop_device(device_id: int) -> int 13torch_npu.npu.stop_device(device_id: int) -> int
14```14```
15 15 
@@ -18,6 +18,7 @@ torch_npu.npu.stop_device(device_id: int) -> int
18**device_id**(`int`):需要处理的device id,确保是一个有效的device。18**device_id**(`int`):需要处理的device id,确保是一个有效的device。
19 19 
20## 返回值说明20## 返回值说明
21+ 
21`int`22`int`
22 23 
23返回值为`int`,代表执行结果,0表示执行成功,1表示执行失败。24返回值为`int`,代表执行结果,0表示执行成功,1表示执行失败。
@@ -30,4 +31,3 @@ torch_npu.npu.stop_device(device_id: int) -> int
30>>> torch.npu.set_device(0) 31>>> torch.npu.set_device(0)
31>>> torch_npu.npu.stop_device(0)32>>> torch_npu.npu.stop_device(0)
32 ```33 ```
33- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-stress_detect.md+4-6
@@ -7,14 +7,13 @@
7|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9 9 
10- 
11## 功能说明10## 功能说明
12 11 
13提供精度在线检测接口,供模型调用。主要通过`StressDetect`接口实现,该接口会对硬件做压力测试检测是否存在静默精度问题。12提供精度在线检测接口,供模型调用。主要通过`StressDetect`接口实现,该接口会对硬件做压力测试检测是否存在静默精度问题。
14 13 
15## 函数原型14## 函数原型
16 15 
17-```16+```python
18torch_npu.npu.stress_detect(detect_type="aic")17torch_npu.npu.stress_detect(detect_type="aic")
19```18```
20 19 
@@ -25,7 +24,6 @@ torch_npu.npu.stress_detect(detect_type="aic")
25> [!NOTE] 24> [!NOTE]
26> 当`detect_type`配置为hccs时,首先基于全局通信域创建本机所有卡的子通信域,然后对该子通信域进行HCCS链路压测。25> 当`detect_type`配置为hccs时,首先基于全局通信域创建本机所有卡的子通信域,然后对该子通信域进行HCCS链路压测。
27 26 
28- 
29## 返回值说明27## 返回值说明
30 28 
31- 接口返回值为`int`,代表错误类型,含义如下所示:29- 接口返回值为`int`,代表错误类型,含义如下所示:
@@ -37,9 +35,11 @@ torch_npu.npu.stress_detect(detect_type="aic")
37 - 2:在线精度检测不通过,硬件故障。35 - 2:在线精度检测不通过,硬件故障。
38 36 
39- 若报如下异常,则表示电压恢复失败,需参见[LINK](https://www.hiascend.com/document/detail/zh/canncommercial/850/maintenref/troubleshooting/troubleshooting_0505.html)手动恢复电压或reboot。37- 若报如下异常,则表示电压恢复失败,需参见[LINK](https://www.hiascend.com/document/detail/zh/canncommercial/850/maintenref/troubleshooting/troubleshooting_0505.html)手动恢复电压或reboot。
40- ```38+ 
39+ ```shell
41 Stress detect error. Error code is 574007. Error message is Voltage recovery failed.40 Stress detect error. Error code is 574007. Error message is Voltage recovery failed.
42 ```41 ```
42+ 
43## 约束说明43## 约束说明
44 44 
45- 精度在线检测的使用需要修改用户的模型训练脚本,建议在训练开始前、结束后以及两个step之间调用,同时需要预留10G大小的内存供压测接口使用。45- 精度在线检测的使用需要修改用户的模型训练脚本,建议在训练开始前、结束后以及两个step之间调用,同时需要预留10G大小的内存供压测接口使用。
@@ -47,7 +47,6 @@ torch_npu.npu.stress_detect(detect_type="aic")
47- 精度在线检测用例,不支持在同一节点运行多个训练作业场景下使用,同时调压功能不支持算力切分场景。47- 精度在线检测用例,不支持在同一节点运行多个训练作业场景下使用,同时调压功能不支持算力切分场景。
48- 不建议使用多线程运行在线精度检测用例。48- 不建议使用多线程运行在线精度检测用例。
49 49 
50- 
51## 调用示例50## 调用示例
52 51 
53```python52```python
@@ -119,4 +118,3 @@ except StressDetectionException as e:
119 print(f"Training halted due to: {e}")118 print(f"Training halted due to: {e}")
120 # do something119 # do something
121```120```
122- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-utils-is_support_inf_nan.md+3-4
@@ -1,4 +1,5 @@
1# (beta)torch_npu.npu.utils.is_support_inf_nan1# (beta)torch_npu.npu.utils.is_support_inf_nan
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -14,19 +15,18 @@
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.npu.utils.is_support_inf_nan() -> bool19torch_npu.npu.utils.is_support_inf_nan() -> bool
19```20```
20 21 
21- 
22## 返回值说明22## 返回值说明
23+ 
23`bool`24`bool`
24 25 
25返回值为True时,代表是INF_NAN模式。26返回值为True时,代表是INF_NAN模式。
26 27 
27返回值为False时,代表为饱和模式。28返回值为False时,代表为饱和模式。
28 29 
29- 
30## 调用示例30## 调用示例
31 31 
32```python32```python
@@ -46,4 +46,3 @@ class TestCheckOverFlow(TestCase):
46if __name__ == "__main__":46if __name__ == "__main__":
47 run_tests()47 run_tests()
48```48```
49- 
Mdocs/zh/custom_APIs/torch_npu-npu/(beta)torch_npu-npu-utils-npu_check_overflow.md+3-4
@@ -9,22 +9,22 @@
9|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
10|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
11 11 
12- 
13- 
14## 功能说明12## 功能说明
15 13 
16检测梯度是否溢出。在INF_NAN模式下,检测输入`Tensor`是否溢出;在饱和模式下,通过检查硬件溢出标志位判断是否溢出。14检测梯度是否溢出。在INF_NAN模式下,检测输入`Tensor`是否溢出;在饱和模式下,通过检查硬件溢出标志位判断是否溢出。
17 15 
18## 函数原型16## 函数原型
19 17 
20-```18+```python
21torch_npu.npu.utils.npu_check_overflow(grad) -> bool19torch_npu.npu.utils.npu_check_overflow(grad) -> bool
22```20```
21+ 
23## 参数说明22## 参数说明
24 23 
25**grad**`Tensor``float`):在INF_NAN模式下判断输入中是否有`inf`或`nan`;饱和模式下,忽略输入,检查硬件溢出标志位。24**grad**`Tensor``float`):在INF_NAN模式下判断输入中是否有`inf`或`nan`;饱和模式下,忽略输入,检查硬件溢出标志位。
26 25 
27## 返回值说明26## 返回值说明
27+ 
28`bool`28`bool`
29 29 
30True溢出,False未溢出。30True溢出,False未溢出。
@@ -50,4 +50,3 @@ class TestCheckOverFlow(TestCase):
50if __name__ == "__main__":50if __name__ == "__main__":
51 run_tests()51 run_tests()
52```52```
53- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedAdadelta.md+1-6
@@ -16,11 +16,10 @@ Adadelta的功能和原理可参考[Adadelta](https://pytorch.org/docs/stable/ge
16 16 
17## 函数原型17## 函数原型
18 18 
19-```19+```python
20class torch_npu.optim.NpuFusedAdadelta(params, lr=1.0, rho=0.9, eps=1e-6, weight_decay=0)20class torch_npu.optim.NpuFusedAdadelta(params, lr=1.0, rho=0.9, eps=1e-6, weight_decay=0)
21```21```
22 22 
23- 
24## 参数说明23## 参数说明
25 24 
26- **params** (`iterable`):必选参数,模型参数或模型参数组。25- **params** (`iterable`):必选参数,模型参数或模型参数组。
@@ -29,12 +28,10 @@ class torch_npu.optim.NpuFusedAdadelta(params, lr=1.0, rho=0.9, eps=1e-6, weight
29- **eps** (`float`):可选参数,分母防止除0项,提高数值稳定性,默认值为1e-6。`eps`小于0时,系统会抛出“ValueError”异常信息。28- **eps** (`float`):可选参数,分母防止除0项,提高数值稳定性,默认值为1e-6。`eps`小于0时,系统会抛出“ValueError”异常信息。
30- **weight_decay** (`float`):可选参数,权重衰减,默认值为0。`weight_decay`小于0时,系统会抛出“ValueError”异常信息。29- **weight_decay** (`float`):可选参数,权重衰减,默认值为0。`weight_decay`小于0时,系统会抛出“ValueError”异常信息。
31 30 
32- 
33## 返回值说明31## 返回值说明
34 32 
35类型为`NpuFusedAdadelta`的对象。33类型为`NpuFusedAdadelta`的对象。
36 34 
37- 
38## 约束说明35## 约束说明
39 36 
40`NpuFusedAdadelta`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:37`NpuFusedAdadelta`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:
@@ -45,7 +42,6 @@ class torch_npu.optim.NpuFusedAdadelta(params, lr=1.0, rho=0.9, eps=1e-6, weight
45 42 
46对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedAdadelta`可正常工作。43对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedAdadelta`可正常工作。
47 44 
48- 
49## 调用示例45## 调用示例
50 46 
51```python47```python
@@ -77,4 +73,3 @@ fused_opt = NpuFusedAdadelta(params, **opt_kwargs)
77with torch.no_grad():73with torch.no_grad():
78 fused_opt.step()74 fused_opt.step()
79```75```
80- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedAdam.md+1-4
@@ -16,7 +16,7 @@ Adam的功能和原理可参考[Adam](https://pytorch.org/docs/stable/generated/
16 16 
17## 函数原型17## 函数原型
18 18 
19-```19+```python
20class torch_npu.optim.NpuFusedAdam(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0, amsgrad=False)20class torch_npu.optim.NpuFusedAdam(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0, amsgrad=False)
21```21```
22 22 
@@ -29,12 +29,10 @@ class torch_npu.optim.NpuFusedAdam(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8
29- **weight_decay** (`float`):可选参数,权重衰减,默认值为0。`weight_decay`小于0时,打印“ValueError”异常信息。29- **weight_decay** (`float`):可选参数,权重衰减,默认值为0。`weight_decay`小于0时,打印“ValueError”异常信息。
30- **amsgrad** (`bool`):可选参数,是否使用AMSGrad,默认值为False。30- **amsgrad** (`bool`):可选参数,是否使用AMSGrad,默认值为False。
31 31 
32- 
33## 返回值说明32## 返回值说明
34 33 
35类型为`NpuFusedAdam`的对象。34类型为`NpuFusedAdam`的对象。
36 35 
37- 
38## 约束说明36## 约束说明
39 37 
40`NpuFusedAdam`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:38`NpuFusedAdam`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:
@@ -76,4 +74,3 @@ fused_opt = NpuFusedAdam(params, **opt_kwargs)
76with torch.no_grad():74with torch.no_grad():
77 fused_opt.step()75 fused_opt.step()
78```76```
79- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedAdamP.md+1-3
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18class torch_npu.optim.NpuFusedAdamP(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0, delta=0.1, wd_ratio=0.1, nesterov=False)18class torch_npu.optim.NpuFusedAdamP(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0, delta=0.1, wd_ratio=0.1, nesterov=False)
19```19```
20 20 
@@ -29,7 +29,6 @@ class torch_npu.optim.NpuFusedAdamP(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-
29- **wd_ratio** (`float`):权重衰减动态调整速率,默认值为0.1。29- **wd_ratio** (`float`):权重衰减动态调整速率,默认值为0.1。
30- **nesterov** (`bool`):使用Nesterov动量,默认值为False。30- **nesterov** (`bool`):使用Nesterov动量,默认值为False。
31 31 
32- 
33## 返回值说明32## 返回值说明
34 33 
35类型为`NpuFusedAdamP`的对象。34类型为`NpuFusedAdamP`的对象。
@@ -75,4 +74,3 @@ fused_opt = NpuFusedAdamP(params, **opt_kwargs)
75with torch.no_grad():74with torch.no_grad():
76 fused_opt.step()75 fused_opt.step()
77```76```
78- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedAdamW.md+1-6
@@ -16,11 +16,10 @@ AdamW的功能和原理可参考[AdamW](https://pytorch.org/docs/stable/generate
16 16 
17## 函数原型17## 函数原型
18 18 
19-```19+```python
20class torch_npu.optim.NpuFusedAdamW(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=1e-2, amsgrad=False)20class torch_npu.optim.NpuFusedAdamW(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=1e-2, amsgrad=False)
21```21```
22 22 
23- 
24## 参数说明23## 参数说明
25 24 
26- **params** (`iterable`):必选参数,模型参数或模型参数组。25- **params** (`iterable`):必选参数,模型参数或模型参数组。
@@ -30,12 +29,10 @@ class torch_npu.optim.NpuFusedAdamW(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-
30- **weight_decay** (`float`):可选参数,权重衰减,默认值为1e-2。`weight_decay`小于0时,系统会抛出“ValueError”异常信息。29- **weight_decay** (`float`):可选参数,权重衰减,默认值为1e-2。`weight_decay`小于0时,系统会抛出“ValueError”异常信息。
31- **amsgrad** (`bool`):可选参数,是否使用参考《[On the Convergence of Adam and Beyond](https://arxiv.org/pdf/1904.09237)》的AMSGrad变种实现,默认值为False。30- **amsgrad** (`bool`):可选参数,是否使用参考《[On the Convergence of Adam and Beyond](https://arxiv.org/pdf/1904.09237)》的AMSGrad变种实现,默认值为False。
32 31 
33- 
34## 返回值说明32## 返回值说明
35 33 
36类型为`NpuFusedAdamW`的对象。34类型为`NpuFusedAdamW`的对象。
37 35 
38- 
39## 约束说明36## 约束说明
40 37 
41`NpuFusedAdamW`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:38`NpuFusedAdamW`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:
@@ -46,7 +43,6 @@ class torch_npu.optim.NpuFusedAdamW(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-
46 43 
47对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedAdamW`可正常工作。44对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedAdamW`可正常工作。
48 45 
49- 
50## 调用示例46## 调用示例
51 47 
52```python48```python
@@ -78,4 +74,3 @@ fused_opt = NpuFusedAdamW(params, **opt_kwargs)
78with torch.no_grad():74with torch.no_grad():
79 fused_opt.step()75 fused_opt.step()
80```76```
81- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedBertAdam.md+1-6
@@ -14,11 +14,10 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18class torch_npu.optim.NpuFusedBertAdam(params, lr=1e-3, warmup=-1, t_total=-1, schedule="warmup_linear", b1=0.9, b2=0.999, e=1e-6, weight_decay=0.01, max_grad_norm=1.0)18class torch_npu.optim.NpuFusedBertAdam(params, lr=1e-3, warmup=-1, t_total=-1, schedule="warmup_linear", b1=0.9, b2=0.999, e=1e-6, weight_decay=0.01, max_grad_norm=1.0)
19```19```
20 20 
21- 
22## 参数说明21## 参数说明
23 22 
24- **params** (`iterable`):必选参数,模型参数或模型参数组。23- **params** (`iterable`):必选参数,模型参数或模型参数组。
@@ -32,12 +31,10 @@ class torch_npu.optim.NpuFusedBertAdam(params, lr=1e-3, warmup=-1, t_total=-1, s
32- **weight_decay** (`float`):可选参数,权重衰减,默认值为0.01。31- **weight_decay** (`float`):可选参数,权重衰减,默认值为0.01。
33- **max_grad_norm** (`float`):最大梯度范围,默认值为1.0,-1表示不做裁剪。32- **max_grad_norm** (`float`):最大梯度范围,默认值为1.0,-1表示不做裁剪。
34 33 
35- 
36## 返回值说明34## 返回值说明
37 35 
38类型为`NpuFusedBertAdam`的对象。36类型为`NpuFusedBertAdam`的对象。
39 37 
40- 
41## 约束说明38## 约束说明
42 39 
43`NpuFusedBertAdam`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:40`NpuFusedBertAdam`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:
@@ -48,7 +45,6 @@ class torch_npu.optim.NpuFusedBertAdam(params, lr=1e-3, warmup=-1, t_total=-1, s
48 45 
49对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedBertAdam`可正常工作。46对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedBertAdam`可正常工作。
50 47 
51- 
52## 调用示例48## 调用示例
53 49 
54```python50```python
@@ -80,4 +76,3 @@ fused_opt = NpuFusedBertAdam(params, **opt_kwargs)
80with torch.no_grad():76with torch.no_grad():
81 fused_opt.step()77 fused_opt.step()
82```78```
83- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedLamb.md+1-6
@@ -14,11 +14,10 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18class torch_npu.optim.NpuFusedLamb(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-6, weight_decay=0, adam=False, use_global_grad_norm=False)18class torch_npu.optim.NpuFusedLamb(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-6, weight_decay=0, adam=False, use_global_grad_norm=False)
19```19```
20 20 
21- 
22## 参数说明21## 参数说明
23 22 
24- **params** (`iterable`):必选参数,模型参数或模型参数组。23- **params** (`iterable`):必选参数,模型参数或模型参数组。
@@ -29,12 +28,10 @@ class torch_npu.optim.NpuFusedLamb(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-6
29- **adam** (`bool`):可选参数,是否通过将trust ratio设置为1,退化为Adam,默认值为False。28- **adam** (`bool`):可选参数,是否通过将trust ratio设置为1,退化为Adam,默认值为False。
30- **use_global_grad_norm** (`bool`):可选参数,是否使用全局梯度正则,默认值为False。29- **use_global_grad_norm** (`bool`):可选参数,是否使用全局梯度正则,默认值为False。
31 30 
32- 
33## 返回值说明31## 返回值说明
34 32 
35类型为`NpuFusedLamb`的对象。33类型为`NpuFusedLamb`的对象。
36 34 
37- 
38## 约束说明35## 约束说明
39 36 
40`NpuFusedLamb`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:37`NpuFusedLamb`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:
@@ -45,7 +42,6 @@ class torch_npu.optim.NpuFusedLamb(params, lr=1e-3, betas=(0.9, 0.999), eps=1e-6
45 42 
46对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedLamb`可正常工作。43对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedLamb`可正常工作。
47 44 
48- 
49## 调用示例45## 调用示例
50 46 
51```python47```python
@@ -77,4 +73,3 @@ fused_opt = NpuFusedLamb(params, **opt_kwargs)
77with torch.no_grad():73with torch.no_grad():
78 fused_opt.step()74 fused_opt.step()
79```75```
80- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedOptimizerBase.md+1-4
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18class torch_npu.optim.NpuFusedOptimizerBase(params, default)18class torch_npu.optim.NpuFusedOptimizerBase(params, default)
19```19```
20 20 
@@ -23,7 +23,6 @@ class torch_npu.optim.NpuFusedOptimizerBase(params, default)
23- **params** (`iterable`):必选参数,模型参数或模型参数组。23- **params** (`iterable`):必选参数,模型参数或模型参数组。
24- **default** (`dict`):包含其他所有参数的字典。24- **default** (`dict`):包含其他所有参数的字典。
25 25 
26- 
27## 返回值说明26## 返回值说明
28 27 
29类型为`NpuFusedOptimizerBase`的对象。28类型为`NpuFusedOptimizerBase`的对象。
@@ -32,7 +31,6 @@ class torch_npu.optim.NpuFusedOptimizerBase(params, default)
32 31 
33`NpuFusedOptimizerBase`为基类,无法单独使用,需通过继承子类实现特定功能的融合优化器。32`NpuFusedOptimizerBase`为基类,无法单独使用,需通过继承子类实现特定功能的融合优化器。
34 33 
35- 
36## 调用示例34## 调用示例
37 35 
38```python36```python
@@ -80,4 +78,3 @@ class NpuFusedSGD(NpuFusedOptimizerBase):
80 for group in self.param_groups: 78 for group in self.param_groups:
81 group.setdefault('nesterov', False)79 group.setdefault('nesterov', False)
82```80```
83- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedRMSprop.md+1-7
@@ -8,7 +8,6 @@
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9|<term>Atlas 训练系列产品</term> | √ |9|<term>Atlas 训练系列产品</term> | √ |
10 10 
11- 
12## 功能说明11## 功能说明
13 12 
14通过张量融合实现的高性能RMSprop优化器,核心功能和`torch.optim.RMSprop`兼容。13通过张量融合实现的高性能RMSprop优化器,核心功能和`torch.optim.RMSprop`兼容。
@@ -17,11 +16,10 @@ RMSprop的功能和原理可参考[RMSprop](https://pytorch.org/docs/stable/gene
17 16 
18## 函数原型17## 函数原型
19 18 
20-```19+```python
21class torch_npu.optim.NpuFusedRMSprop(params, lr=1e-2, alpha=0.99, eps=1e-8, weight_decay=0, momentum=0, centered=False)20class torch_npu.optim.NpuFusedRMSprop(params, lr=1e-2, alpha=0.99, eps=1e-8, weight_decay=0, momentum=0, centered=False)
22```21```
23 22 
24- 
25## 参数说明23## 参数说明
26 24 
27- **params** (`iterable`):必选参数,模型参数或模型参数组。25- **params** (`iterable`):必选参数,模型参数或模型参数组。
@@ -32,12 +30,10 @@ class torch_npu.optim.NpuFusedRMSprop(params, lr=1e-2, alpha=0.99, eps=1e-8, wei
32- **momentum** (`float`):可选参数,动量因子,默认值为0。`momentum`的值小于0时,打印“ValueError”异常信息。30- **momentum** (`float`):可选参数,动量因子,默认值为0。`momentum`的值小于0时,打印“ValueError”异常信息。
33- **centered** (`bool`):可选参数,计算中心RMSProp,梯度将被方差的估计值归一化,默认值为False。31- **centered** (`bool`):可选参数,计算中心RMSProp,梯度将被方差的估计值归一化,默认值为False。
34 32 
35- 
36## 返回值说明33## 返回值说明
37 34 
38类型为`NpuFusedRMSprop`的对象。35类型为`NpuFusedRMSprop`的对象。
39 36 
40- 
41## 约束说明37## 约束说明
42 38 
43`NpuFusedRMSprop`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:39`NpuFusedRMSprop`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:
@@ -48,7 +44,6 @@ class torch_npu.optim.NpuFusedRMSprop(params, lr=1e-2, alpha=0.99, eps=1e-8, wei
48 44 
49对模型参数对象进行inplace计算,或读取参数的值时,`NpuFusedRMSprop`都可以正常工作。45对模型参数对象进行inplace计算,或读取参数的值时,`NpuFusedRMSprop`都可以正常工作。
50 46 
51- 
52## 调用示例47## 调用示例
53 48 
54```python49```python
@@ -80,4 +75,3 @@ fused_opt = NpuFusedRMSprop(params, **opt_kwargs)
80with torch.no_grad():75with torch.no_grad():
81 fused_opt.step()76 fused_opt.step()
82```77```
83- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedRMSpropTF.md+1-5
@@ -14,11 +14,10 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18class torch_npu.optim.NpuFusedRMSpropTF(params, lr=1e-2, alpha=0.9, eps=1e-10, weight_decay=0, momentum=0., centered=False, decoupled_decay=False, lr_in_momentum=True)18class torch_npu.optim.NpuFusedRMSpropTF(params, lr=1e-2, alpha=0.9, eps=1e-10, weight_decay=0, momentum=0., centered=False, decoupled_decay=False, lr_in_momentum=True)
19```19```
20 20 
21- 
22## 参数说明21## 参数说明
23 22 
24- **params** (`iterable`):必选参数,模型参数或模型参数组。23- **params** (`iterable`):必选参数,模型参数或模型参数组。
@@ -31,12 +30,10 @@ class torch_npu.optim.NpuFusedRMSpropTF(params, lr=1e-2, alpha=0.9, eps=1e-10, w
31- **decoupled_decay** (`bool`):可选参数,权重衰减仅作用于参数,默认值为False。30- **decoupled_decay** (`bool`):可选参数,权重衰减仅作用于参数,默认值为False。
32- **lr_in_momentum** (`bool`):可选参数,计算动量buffer时使用lr,默认值为True。31- **lr_in_momentum** (`bool`):可选参数,计算动量buffer时使用lr,默认值为True。
33 32 
34- 
35## 返回值说明33## 返回值说明
36 34 
37类型为`NpuFusedRMSpropTF`的对象。35类型为`NpuFusedRMSpropTF`的对象。
38 36 
39- 
40## 约束说明37## 约束说明
41 38 
42`NpuFusedRMSpropTF`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:39`NpuFusedRMSpropTF`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:
@@ -78,4 +75,3 @@ fused_opt = NpuFusedRMSpropTF(params, **opt_kwargs)
78with torch.no_grad():75with torch.no_grad():
79 fused_opt.step()76 fused_opt.step()
80```77```
81- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim-NpuFusedSGD.md+1-5
@@ -16,7 +16,7 @@ SGD的功能和原理可参考[SGD](https://pytorch.org/docs/stable/generated/to
16 16 
17## 函数原型17## 函数原型
18 18 
19-```19+```python
20class torch_npu.optim.NpuFusedSGD(params, lr, momentum=0, dampening=0, weight_decay=0, nesterov=False)20class torch_npu.optim.NpuFusedSGD(params, lr, momentum=0, dampening=0, weight_decay=0, nesterov=False)
21```21```
22 22 
@@ -29,12 +29,10 @@ class torch_npu.optim.NpuFusedSGD(params, lr, momentum=0, dampening=0, weight_de
29- **weight_decay** (`float`):可选参数,权重衰减,默认值为0。`weight_decay`小于0时,系统会抛出“ValueError”异常信息。29- **weight_decay** (`float`):可选参数,权重衰减,默认值为0。`weight_decay`小于0时,系统会抛出“ValueError”异常信息。
30- **nesterov** (`bool`):可选参数,是否使用Nesterov动量,默认值为False。`nesterov`为True,同时`momentum`小于0或者`dampening`不等于0,系统会抛出“TypeError”异常信息。30- **nesterov** (`bool`):可选参数,是否使用Nesterov动量,默认值为False。`nesterov`为True,同时`momentum`小于0或者`dampening`不等于0,系统会抛出“TypeError”异常信息。
31 31 
32- 
33## 返回值说明32## 返回值说明
34 33 
35类型为`NpuFusedSGD`的对象。34类型为`NpuFusedSGD`的对象。
36 35 
37- 
38## 约束说明36## 约束说明
39 37 
40`NpuFusedSGD`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:38`NpuFusedSGD`的实现机制要求`params`中的每一个模型参数对象在使用过程中不能被重新申请,否则将导致无法预料的结果。引起模型参数对象被重新申请的操作包括但不限于:
@@ -45,7 +43,6 @@ class torch_npu.optim.NpuFusedSGD(params, lr, momentum=0, dampening=0, weight_de
45 43 
46对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedSGD`可正常工作。44对模型参数对象进行inplace计算,或者读取参数的值,`NpuFusedSGD`可正常工作。
47 45 
48- 
49## 调用示例46## 调用示例
50 47 
51```python48```python
@@ -77,4 +74,3 @@ fused_opt = NpuFusedSGD(params, **opt_kwargs)
77with torch.no_grad():74with torch.no_grad():
78 fused_opt.step()75 fused_opt.step()
79```76```
80- 
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim.md+1-1
@@ -1 +1 @@
1-# torch_npu.optim1+# torch_npu.optim
Mdocs/zh/custom_APIs/torch_npu-optim/torch_npu-optim_list.md+0-1
@@ -63,4 +63,3 @@
63</tr>63</tr>
64</tbody>64</tbody>
65</table>65</table>
66- 
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-AiCMetrics.md+2-2
@@ -14,7 +14,7 @@ AI Core的性能指标采集项,Enum类型。用于作为_ExperimentalConfig
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.AiCMetrics18torch_npu.profiler.AiCMetrics
19```19```
20 20 
@@ -57,4 +57,4 @@ with torch_npu.profiler.profile(
57 for step in range(steps): # 训练函数57 for step in range(steps): # 训练函数
58 train_one_step() # 训练函数58 train_one_step() # 训练函数
59 prof.step()59 prof.step()
60-```60+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-ExportType.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.ExportType18torch_npu.profiler.ExportType
19```19```
20 20 
@@ -50,4 +50,4 @@ with torch_npu.profiler.profile(
50 for step in range(steps): # 训练函数50 for step in range(steps): # 训练函数
51 train_one_step() # 训练函数51 train_one_step() # 训练函数
52 prof.step()52 prof.step()
53-```53+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-ProfilerAction.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.ProfilerAction18torch_npu.profiler.ProfilerAction
19```19```
20 20 
@@ -45,4 +45,4 @@ with torch_npu.profiler.profile(
45 for step in range(steps): # 训练函数45 for step in range(steps): # 训练函数
46 train_one_step() # 训练函数46 train_one_step() # 训练函数
47 prof.step()47 prof.step()
48-```48+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-ProfilerActivity.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.ProfilerActivity18torch_npu.profiler.ProfilerActivity
19```19```
20 20 
@@ -49,4 +49,4 @@ with torch_npu.profiler.profile(
49 for step in range(steps): # 训练函数49 for step in range(steps): # 训练函数
50 train_one_step() # 训练函数50 train_one_step() # 训练函数
51 prof.step()51 prof.step()
52-```52+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-ProfilerLevel.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.ProfilerLevel18torch_npu.profiler.ProfilerLevel
19```19```
20 20 
@@ -51,4 +51,4 @@ with torch_npu.profiler.profile(
51 for step in range(steps): # 训练函数51 for step in range(steps): # 训练函数
52 train_one_step() # 训练函数52 train_one_step() # 训练函数
53 prof.step()53 prof.step()
54-```54+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-_ExperimentalConfig.md+3-3
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler._ExperimentalConfig(export_type=[torch_npu.profiler.ExportType.Text], profiler_level=torch_npu.profiler.ProfilerLevel.Level0, mstx=False, mstx_domain_include=[], mstx_domain_exclude=[], aic_metrics=torch_npu.profiler.AiCMetrics.AiCoreNone, l2_cache=False, op_attr=False, data_simplification=True, record_op_args=False, gc_detect_threshold=None, host_sys=[], sys_io=False, sys_interconnection=False)18torch_npu.profiler._ExperimentalConfig(export_type=[torch_npu.profiler.ExportType.Text], profiler_level=torch_npu.profiler.ProfilerLevel.Level0, mstx=False, mstx_domain_include=[], mstx_domain_exclude=[], aic_metrics=torch_npu.profiler.AiCMetrics.AiCoreNone, l2_cache=False, op_attr=False, data_simplification=True, record_op_args=False, gc_detect_threshold=None, host_sys=[], sys_io=False, sys_interconnection=False)
19```19```
20 20 
@@ -35,7 +35,7 @@ torch_npu.profiler._ExperimentalConfig(export_type=[torch_npu.profiler.ExportTyp
35 35 
36- **mstx_domain_include** (`list`):可选参数,输出需要的domain数据。调用torch_npu.npu.mstx系列打点接口,使用默认domain或指定domain进行打点时,可选择只输出本参数配置的domain数据。36- **mstx_domain_include** (`list`):可选参数,输出需要的domain数据。调用torch_npu.npu.mstx系列打点接口,使用默认domain或指定domain进行打点时,可选择只输出本参数配置的domain数据。
37 37 
38- domain名称为用户调用[torch_npu.npu.mstx](torch_npu-npu-mstx.md)系列接口传入的domain或默认domain('default'),domain名称使用list类型输入。38+ domain名称为用户调用[torch_npu.npu.mstx](../torch_npu-npu/torch_npu-npu-mstx.md)系列接口传入的domain或默认domain('default'),domain名称使用list类型输入。
39 39 
40 与mstx_domain_exclude参数互斥,若同时配置,则只有mstx_domain_include生效。40 与mstx_domain_exclude参数互斥,若同时配置,则只有mstx_domain_include生效。
41 41 
@@ -144,4 +144,4 @@ with torch_npu.profiler.profile(
144 for step in range(steps): # 训练函数144 for step in range(steps): # 训练函数
145 train_one_step() # 训练函数145 train_one_step() # 训练函数
146 prof.step()146 prof.step()
147-```147+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-_KinetoProfile.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler._KinetoProfile(activities=None, record_shapes=False, profile_memory=False, with_stack=False, with_flops=False, with_modules=False, experimental_config=None)18torch_npu.profiler._KinetoProfile(activities=None, record_shapes=False, profile_memory=False, with_stack=False, with_flops=False, with_modules=False, experimental_config=None)
19```19```
20 20 
@@ -83,4 +83,4 @@ for epoch in range(epochs):
83 if epoch == 1:83 if epoch == 1:
84 prof.stop()84 prof.stop()
85prof.export_chrome_trace("result_dir/trace.json")85prof.export_chrome_trace("result_dir/trace.json")
86-```86+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-dynamic_profile-init.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.dynamic_profile.init(path: str)18torch_npu.profiler.dynamic_profile.init(path: str)
19```19```
20 20 
@@ -42,4 +42,4 @@ for step in steps:
42 train_one_step()42 train_one_step()
43 # 划分step43 # 划分step
44 dp.step()44 dp.step()
45-```45+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-dynamic_profile-start.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.dynamic_profile.start(config_path: str = None)18torch_npu.profiler.dynamic_profile.start(config_path: str = None)
19```19```
20 20 
@@ -45,4 +45,4 @@ for step in steps:
45 train_one_step()45 train_one_step()
46 # 划分step,需要进行profile的代码需在dp.start()接口和dp.step()接口之间46 # 划分step,需要进行profile的代码需在dp.start()接口和dp.step()接口之间
47 dp.step()47 dp.step()
48-```48+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-dynamic_profile-step.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.dynamic_profile.step()18torch_npu.profiler.dynamic_profile.step()
19```19```
20 20 
@@ -40,4 +40,4 @@ for step in steps:
40 train_one_step()40 train_one_step()
41 # 划分step41 # 划分step
42 dp.step()42 dp.step()
43-```43+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-profile-disable_profiler_in_child_thread.md+1-2
@@ -11,7 +11,6 @@
11| <term>Atlas 推理系列产品</term> | √ |11| <term>Atlas 推理系列产品</term> | √ |
12| <term>Atlas 训练系列产品</term> | √ |12| <term>Atlas 训练系列产品</term> | √ |
13 13 
14- 
15## 功能说明14## 功能说明
16 15 
17注销Profiler采集回调函数。16注销Profiler采集回调函数。
@@ -34,4 +33,4 @@ torch_npu.profiler.profile.disable_profiler_in_child_thread()
34 33 
35## 调用示例34## 调用示例
36 35 
37-见[torch_npu.profiler.profile.enable_profiler_in_child_thread](./torch_npu-profiler-profile-enable_profiler_in_child_thread.md)。36+见[torch_npu.profiler.profile.enable_profiler_in_child_thread](./torch_npu-profiler-profile-enable_profiler_in_child_thread.md)。
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-profile-enable_profiler_in_child_thread.md+1-2
@@ -11,7 +11,6 @@
11| <term>Atlas 推理系列产品</term> | √ |11| <term>Atlas 推理系列产品</term> | √ |
12| <term>Atlas 训练系列产品</term> | √ |12| <term>Atlas 训练系列产品</term> | √ |
13 13 
14- 
15## 功能说明14## 功能说明
16 15 
17注册Profiler采集回调函数,采集用户子线程下发的torch算子等框架侧数据。该参数中可另外配置[torch_npu.profiler.profile](./torch_npu-profiler-profile.md)的参数(包括record_shapes、profile_memory、with_stack、with_flops、with_modules),作为Profiler子线程的采集配置。16注册Profiler采集回调函数,采集用户子线程下发的torch算子等框架侧数据。该参数中可另外配置[torch_npu.profiler.profile](./torch_npu-profiler-profile.md)的参数(包括record_shapes、profile_memory、with_stack、with_flops、with_modules),作为Profiler子线程的采集配置。
@@ -136,4 +135,4 @@ if __name__ == "__main__":
136 t.join()135 t.join()
137 136 
138 prof.stop()137 prof.stop()
139-```138+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-profile.md+2-3
@@ -8,14 +8,13 @@
8| <term>Atlas A2 训练系列产品</term> | √ |8| <term>Atlas A2 训练系列产品</term> | √ |
9| <term>Atlas 训练系列产品</term> | √ |9| <term>Atlas 训练系列产品</term> | √ |
10 10 
11- 
12## 功能说明11## 功能说明
13 12 
14提供PyTorch训练过程中的性能数据采集功能。13提供PyTorch训练过程中的性能数据采集功能。
15 14 
16## 函数原型15## 函数原型
17 16 
18-```17+```python
19torch_npu.profiler.profile(activities=None, schedule=None, on_trace_ready=None, record_shapes=False, profile_memory=False, with_stack=False, with_modules, with_flops=False, experimental_config=None)18torch_npu.profiler.profile(activities=None, schedule=None, on_trace_ready=None, record_shapes=False, profile_memory=False, with_stack=False, with_modules, with_flops=False, experimental_config=None)
20```19```
21 20 
@@ -126,4 +125,4 @@ with torch_npu.profiler.profile(
126 for step in range(steps): # 训练函数125 for step in range(steps): # 训练函数
127 train_one_step() # 训练函数126 train_one_step() # 训练函数
128 prof.step()127 prof.step()
129-```128+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-profiler-analyse.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.profiler.analyse(profiler_path="", max_process_number=max_process_number, export_type=export_type)18torch_npu.profiler.profiler.analyse(profiler_path="", max_process_number=max_process_number, export_type=export_type)
19```19```
20 20 
@@ -46,4 +46,4 @@ from torch_npu.profiler.profiler import analyse
46 46 
47if __name__ == "__main__":47if __name__ == "__main__":
48 analyse(profiler_path="./result_data", max_process_number=max_process_number)48 analyse(profiler_path="./result_data", max_process_number=max_process_number)
49-```49+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-schedule.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.schedule (wait, active, warmup = 0, repeat = 0, skip_first = 0)18torch_npu.profiler.schedule (wait, active, warmup = 0, repeat = 0, skip_first = 0)
19```19```
20 20 
@@ -66,4 +66,4 @@ with torch_npu.profiler.profile(
66 for _ in range(9):66 for _ in range(9):
67 train_one_step()67 train_one_step()
68 prof.step() # 通知profiler完成一个step68 prof.step() # 通知profiler完成一个step
69-```69+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-supported_activities.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.supported_activities()18torch_npu.profiler.supported_activities()
19```19```
20 20 
@@ -33,4 +33,4 @@ import torch_npu
33...33...
34 34 
35torch_npu.profiler.supported_activities()35torch_npu.profiler.supported_activities()
36-```36+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-supported_ai_core_metrics.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.supported_ai_core_metrics()18torch_npu.profiler.supported_ai_core_metrics()
19```19```
20 20 
@@ -33,4 +33,4 @@ import torch_npu
33...33...
34 34 
35torch_npu.profiler.supported_ai_core_metrics()35torch_npu.profiler.supported_ai_core_metrics()
36-```36+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-supported_export_type.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.supported_export_type()18torch_npu.profiler.supported_export_type()
19```19```
20 20 
@@ -33,4 +33,4 @@ import torch_npu
33...33...
34 34 
35torch_npu.profiler.supported_export_type()35torch_npu.profiler.supported_export_type()
36-```36+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-supported_profiler_level.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.supported_profiler_level()18torch_npu.profiler.supported_profiler_level()
19```19```
20 20 
@@ -33,4 +33,4 @@ import torch_npu
33...33...
34 34 
35torch_npu.profiler.supported_profiler_level()35torch_npu.profiler.supported_profiler_level()
36-```36+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler-tensorboard_trace_handler.md+2-2
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch_npu.profiler.tensorboard_trace_handler(dir_name=None, worker_name=None, analyse_flag=True, async_mode=False)18torch_npu.profiler.tensorboard_trace_handler(dir_name=None, worker_name=None, analyse_flag=True, async_mode=False)
19```19```
20 20 
@@ -60,4 +60,4 @@ with torch_npu.profiler.profile(
60 for step in range(steps): # 训练函数60 for step in range(steps): # 训练函数
61 train_one_step() # 训练函数61 train_one_step() # 训练函数
62 prof.step()62 prof.step()
63-```63+```
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler.md+1-1
@@ -1 +1 @@
1-# torch_npu.profiler1+# torch_npu.profiler
Mdocs/zh/custom_APIs/torch_npu-profiler/torch_npu-profiler_list.md+5-5
@@ -18,9 +18,9 @@
18| [torch_npu.profiler.supported_profiler_level](./torch_npu-profiler-supported_profiler_level.md) | 查询当前支持的torch_npu.profiler.ProfilerLevel级别。 |18| [torch_npu.profiler.supported_profiler_level](./torch_npu-profiler-supported_profiler_level.md) | 查询当前支持的torch_npu.profiler.ProfilerLevel级别。 |
19| [torch_npu.profiler.supported_ai_core_metrics](./torch_npu-profiler-supported_ai_core_metrics.md) | 查询当前支持的torch_npu.profiler. AiCMetrics的AI Core性能指标采集项。 |19| [torch_npu.profiler.supported_ai_core_metrics](./torch_npu-profiler-supported_ai_core_metrics.md) | 查询当前支持的torch_npu.profiler. AiCMetrics的AI Core性能指标采集项。 |
20| [torch_npu.profiler.supported_export_type](./torch_npu-profiler-supported_export_type.md) | 查询当前支持的torch_npu.profiler.ExportType的性能数据结果文件类型。 |20| [torch_npu.profiler.supported_export_type](./torch_npu-profiler-supported_export_type.md) | 查询当前支持的torch_npu.profiler.ExportType的性能数据结果文件类型。 |
21-| [torch_npu.profiler.dynamic_profile.init](./torch_npu-profiler-dynamic_profile.init.md) | 初始化dynamic_profile动态采集。 |21+| [torch_npu.profiler.dynamic_profile.init](./torch_npu-profiler-dynamic_profile-init.md) | 初始化dynamic_profile动态采集。 |
22-| [torch_npu.profiler.dynamic_profile.step](./torch_npu-profiler-dynamic_profile.step.md) | dynamic_profile动态采集划分step。 |22+| [torch_npu.profiler.dynamic_profile.step](./torch_npu-profiler-dynamic_profile-step.md) | dynamic_profile动态采集划分step。 |
23-| [torch_npu.profiler.dynamic_profile.start](./torch_npu-profiler-dynamic_profile.start.md) | 触发一次dynamic_profile动态采集。 |23+| [torch_npu.profiler.dynamic_profile.start](./torch_npu-profiler-dynamic_profile-start.md) | 触发一次dynamic_profile动态采集。 |
24-| [torch_npu.profiler.profiler.analyse](./torch_npu-profiler-profiler.analyse.md) | Ascend PyTorch Profiler性能数据离线解析。 |24+| [torch_npu.profiler.profiler.analyse](./torch_npu-profiler-profiler-analyse.md) | Ascend PyTorch Profiler性能数据离线解析。 |
25| [torch_npu.profiler.profile.enable_profiler_in_child_thread](./torch_npu-profiler-profile-enable_profiler_in_child_thread.md) | 注册Profiler采集回调函数。 |25| [torch_npu.profiler.profile.enable_profiler_in_child_thread](./torch_npu-profiler-profile-enable_profiler_in_child_thread.md) | 注册Profiler采集回调函数。 |
26-| [torch_npu.profiler.profile.disable_profiler_in_child_thread](./torch_npu-profiler-profile-disable_profiler_in_child_thread.md) | 注销Profiler采集回调函数。 |26+| [torch_npu.profiler.profile.disable_profiler_in_child_thread](./torch_npu-profiler-profile-disable_profiler_in_child_thread.md) | 注销Profiler采集回调函数。 |
Mdocs/zh/custom_APIs/torch_npu-utils/torch_npu-utils-get_cann_version.md+4-3
@@ -1,4 +1,5 @@
1# torch_npu.utils.get_cann_version1# torch_npu.utils.get_cann_version
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -13,7 +14,7 @@
13 14 
14## 函数原型15## 函数原型
15 16 
16-```17+```python
17torch_npu.utils.get_cann_version(module="CANN")18torch_npu.utils.get_cann_version(module="CANN")
18```19```
19 20 
@@ -39,8 +40,8 @@ torch_npu.utils.get_cann_version(module="CANN")
39 40 
40- "DRIVER":驱动。41- "DRIVER":驱动。
41 42 
42- 
43## 返回值说明43## 返回值说明
44+ 
44`str`45`str`
45 46 
46代表具体组件的版本号。47代表具体组件的版本号。
@@ -61,4 +62,4 @@ torch_npu.utils.get_cann_version(module="CANN")
61>>> version = get_cann_version(module="CANN")62>>> version = get_cann_version(module="CANN")
62>>> version63>>> version
63'8.3.RC1'64'8.3.RC1'
64-```65+```
Mdocs/zh/custom_APIs/torch_npu-utils/torch_npu-utils.md+1-1
@@ -1 +1 @@
1-# torch_npu.utils1+# torch_npu.utils
Mdocs/zh/custom_APIs/torch_npu-utils/torch_npu-utils.reset_thread_affinity.md+3-5
@@ -1,4 +1,5 @@
1# torch_npu.utils.reset_thread_affinity1# torch_npu.utils.reset_thread_affinity
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -10,14 +11,13 @@
10|<term>Atlas 推理系列产品</term> | √ |11|<term>Atlas 推理系列产品</term> | √ |
11|<term>Atlas 训练系列产品</term> | √ |12|<term>Atlas 训练系列产品</term> | √ |
12 13 
13- 
14## 功能说明14## 功能说明
15 15 
16恢复当前线程的绑核区间为主线程。16恢复当前线程的绑核区间为主线程。
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.utils.reset_thread_affinity()21torch_npu.utils.reset_thread_affinity()
22```22```
23 23 
@@ -25,18 +25,16 @@ torch_npu.utils.reset_thread_affinity()
25 25 
2626
27 27 
28- 
29## 返回值说明28## 返回值说明
29+ 
3030
31 31 
32## 约束说明32## 约束说明
33 33 
34该接口需要环境变量`CPU_AFFINITY_CONF`的mode设置为1或2时才生效,一般在拉起子线程的位置后使用,恢复当前线程的绑核区间为主线程区间。推荐和[torch_npu.utils.set_thread_affinity](torch_npu-utils.set_thread_affinity.md)配套使用。34该接口需要环境变量`CPU_AFFINITY_CONF`的mode设置为1或2时才生效,一般在拉起子线程的位置后使用,恢复当前线程的绑核区间为主线程区间。推荐和[torch_npu.utils.set_thread_affinity](torch_npu-utils.set_thread_affinity.md)配套使用。
35 35 
36- 
37## 调用示例36## 调用示例
38 37 
39- 
40```python38```python
41>>> import torch_npu39>>> import torch_npu
42>>> import threading40>>> import threading
Mdocs/zh/custom_APIs/torch_npu-utils/torch_npu-utils.set_thread_affinity.md+3-5
@@ -1,4 +1,5 @@
1# torch_npu.utils.set_thread_affinity1# torch_npu.utils.set_thread_affinity
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -10,14 +11,13 @@
10|<term>Atlas 推理系列产品</term> | √ |11|<term>Atlas 推理系列产品</term> | √ |
11|<term>Atlas 训练系列产品</term> | √ |12|<term>Atlas 训练系列产品</term> | √ |
12 13 
13- 
14## 功能说明14## 功能说明
15 15 
16设置当前线程的绑核区间。16设置当前线程的绑核区间。
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.utils.set_thread_affinity(core_range: List[int] = None)21torch_npu.utils.set_thread_affinity(core_range: List[int] = None)
22```22```
23 23 
@@ -25,18 +25,16 @@ torch_npu.utils.set_thread_affinity(core_range: List[int] = None)
25 25 
26 **core_range** (`List[int]`):可选参数,表示用户期望设置的当前线程绑核区间。默认值为`None`,表示将当前线程作为非主要线程进行自动绑核。26 **core_range** (`List[int]`):可选参数,表示用户期望设置的当前线程绑核区间。默认值为`None`,表示将当前线程作为非主要线程进行自动绑核。
27 27 
28- 
29## 返回值说明28## 返回值说明
29+ 
3030
31 31 
32## 约束说明32## 约束说明
33 33 
34该接口需要环境变量`CPU_AFFINITY_CONF`的mode设置为1或2时才生效,一般在拉起子线程的位置前使用,指定子线程的绑核方式或绑核区间。推荐和[torch_npu.utils.reset_thread_affinity](torch_npu-utils.reset_thread_affinity.md)配套使用。34该接口需要环境变量`CPU_AFFINITY_CONF`的mode设置为1或2时才生效,一般在拉起子线程的位置前使用,指定子线程的绑核方式或绑核区间。推荐和[torch_npu.utils.reset_thread_affinity](torch_npu-utils.reset_thread_affinity.md)配套使用。
35 35 
36- 
37## 调用示例36## 调用示例
38 37 
39- 
40```python38```python
41>>> import torch_npu39>>> import torch_npu
42>>> import threading40>>> import threading
Mdocs/zh/custom_APIs/torch_npu-utils/torch_npu-utils_list.md+0-1
@@ -53,4 +53,3 @@
53</tr>53</tr>
54</tbody>54</tbody>
55</table>55</table>
56- 
Mdocs/zh/custom_APIs/torch_npu-utils/(beta)torch_npu-utils-FlopsCounter.md+1-3
@@ -16,7 +16,7 @@ torch_npu\utils\flops_count.py
16 16 
17## 函数原型17## 函数原型
18 18 
19-```19+```python
20torch_npu.utils.FlopsCounter()20torch_npu.utils.FlopsCounter()
21```21```
22 22 
@@ -56,7 +56,6 @@ torch_npu.utils.FlopsCounter()
56 56 
57 获取统计结果。返回列表,包括不含重计算的Flops(recordedCount)和含重计算的Flops(traversedCount),例如[_100, 200_],_100_为不含重计算的Flops(recordedCount),_200_为含重计算的Flops(traversedCount)。57 获取统计结果。返回列表,包括不含重计算的Flops(recordedCount)和含重计算的Flops(traversedCount),例如[_100, 200_],_100_为不含重计算的Flops(recordedCount),_200_为含重计算的Flops(traversedCount)。
58 58 
59- 
60## 调用示例59## 调用示例
61 60 
62```python61```python
@@ -90,4 +89,3 @@ FlopsCounter.stop()
90matmul()89matmul()
91print(f"FlopsCounter.stop():{FlopsCounter.get_flops()}") # 含重计算Flops和不含重计算Flops清0且均不累计90print(f"FlopsCounter.stop():{FlopsCounter.get_flops()}") # 含重计算Flops和不含重计算Flops清0且均不累计
92```91```
93- 
Mdocs/zh/custom_APIs/torch_npu-utils/(beta)torch_npu-utils-get_part_combined_tensor.md+3-1
@@ -1,4 +1,5 @@
1# (beta)torch_npu.utils.get_part_combined_tensor1# (beta)torch_npu.utils.get_part_combined_tensor
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -14,7 +15,7 @@
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.utils.get_part_combined_tensor(combined_tensor, index, size) -> Tensor19torch_npu.utils.get_part_combined_tensor(combined_tensor, index, size) -> Tensor
19```20```
20 21 
@@ -25,6 +26,7 @@ torch_npu.utils.get_part_combined_tensor(combined_tensor, index, size) -> Tensor
25- **size** (`Long`):需获取的局部Tensor的大小。26- **size** (`Long`):需获取的局部Tensor的大小。
26 27 
27## 返回值说明28## 返回值说明
29+ 
28`Tensor`30`Tensor`
29 31 
30代表从融合Tensor中获取的局部Tensor。32代表从融合Tensor中获取的局部Tensor。
Mdocs/zh/custom_APIs/torch_npu-utils/(beta)torch_npu-utils-is_combined_tensor_valid.md+3-2
@@ -1,4 +1,5 @@
1# (beta)torch_npu.utils.is_combined_tensor_valid1# (beta)torch_npu.utils.is_combined_tensor_valid
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -14,7 +15,7 @@
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.utils.is_combined_tensor_valid(combined_tensor, list_of_tensor) -> bool19torch_npu.utils.is_combined_tensor_valid(combined_tensor, list_of_tensor) -> bool
19```20```
20 21 
@@ -24,6 +25,7 @@ torch_npu.utils.is_combined_tensor_valid(combined_tensor, list_of_tensor) -> boo
24- **list_of_tensor** (`List[Tensor]`):需要进行校验的Tensor列表。25- **list_of_tensor** (`List[Tensor]`):需要进行校验的Tensor列表。
25 26 
26## 返回值说明27## 返回值说明
28+ 
27`bool`29`bool`
28 30 
29代表Tensor列表`list_of_tensor`中的Tensor是否全部属于融合Tensor `combined_tensor`31代表Tensor列表`list_of_tensor`中的Tensor是否全部属于融合Tensor `combined_tensor`
@@ -31,4 +33,3 @@ torch_npu.utils.is_combined_tensor_valid(combined_tensor, list_of_tensor) -> boo
31## 约束说明33## 约束说明
32 34 
33融合Tensor `combined_tensor``list_of_tensor`中的Tensor须全部为内存连续的、dtype一致的NPU Tensor。35融合Tensor `combined_tensor``list_of_tensor`中的Tensor须全部为内存连续的、dtype一致的NPU Tensor。
34- 
Mdocs/zh/custom_APIs/torch_npu-utils/(beta)torch_npu-utils-npu_combine_tensors.md+3-1
@@ -1,4 +1,5 @@
1# (beta)torch_npu.utils.npu_combine_tensors1# (beta)torch_npu.utils.npu_combine_tensors
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -14,7 +15,7 @@
14 15 
15## 函数原型16## 函数原型
16 17 
17-```18+```python
18torch_npu.utils.npu_combine_tensors(list_of_tensor, require_copy_value=True) -> Tensor19torch_npu.utils.npu_combine_tensors(list_of_tensor, require_copy_value=True) -> Tensor
19```20```
20 21 
@@ -24,6 +25,7 @@ torch_npu.utils.npu_combine_tensors(list_of_tensor, require_copy_value=True) ->
24- **require_copy_value** (`bool`):默认值为True,是否将原Tensor的值拷贝到新融合Tensor的对应偏移地址。25- **require_copy_value** (`bool`):默认值为True,是否将原Tensor的值拷贝到新融合Tensor的对应偏移地址。
25 26 
26## 返回值说明27## 返回值说明
28+ 
27`Tensor`29`Tensor`
28 30 
29代表融合后的新Tensor。31代表融合后的新Tensor。
Mdocs/zh/custom_APIs/torch_npu-utils/(beta)torch_npu-utils-save_async.md+2-3
@@ -1,4 +1,5 @@
1# (beta)torch\_npu.utils.save\_async1# (beta)torch\_npu.utils.save\_async
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -8,14 +9,13 @@
8|<term>Atlas 推理系列产品</term> | √ |9|<term>Atlas 推理系列产品</term> | √ |
9|<term>Atlas 训练系列产品</term> | √ |10|<term>Atlas 训练系列产品</term> | √ |
10 11 
11- 
12## 功能说明12## 功能说明
13 13 
14异步保存一个对象到一个硬盘文件上。14异步保存一个对象到一个硬盘文件上。
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.utils.save_async(obj, f, pickle_module=pickle, pickle_protocol=DEFAULT_PROTOCOL, _use_new_zipfile_serialization=True, _disable_byteorder_record=False, model=None)19torch_npu.utils.save_async(obj, f, pickle_module=pickle, pickle_protocol=DEFAULT_PROTOCOL, _use_new_zipfile_serialization=True, _disable_byteorder_record=False, model=None)
20```20```
21 21 
@@ -62,4 +62,3 @@ for epoch in range(3):
62 save_path = os.path.join(f"model_{epoch}_{step}.path")62 save_path = os.path.join(f"model_{epoch}_{step}.path")
63 torch_npu.utils.save_async(model, save_path, model=model)63 torch_npu.utils.save_async(model, save_path, model=model)
64```64```
65- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-empty_with_swapped_memory.md+3-4
@@ -1,4 +1,5 @@
1# torch_npu.empty_with_swapped_memory1# torch_npu.empty_with_swapped_memory
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -6,14 +7,13 @@
6|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
7|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
8 9 
9- 
10## 功能说明10## 功能说明
11 11 
12申请一个device信息为NPU且实际内存在host侧的`Tensor`12申请一个device信息为NPU且实际内存在host侧的`Tensor`
13 13 
14## 函数原型14## 函数原型
15 15 
16-```16+```python
17torch_npu.empty_with_swapped_memory(size, dtype=None, device=None) -> Tensor17torch_npu.empty_with_swapped_memory(size, dtype=None, device=None) -> Tensor
18```18```
19 19 
@@ -23,8 +23,8 @@ torch_npu.empty_with_swapped_memory(size, dtype=None, device=None) -> Tensor
23- **dtype** (`torch.dtype`):可选参数,表示生成Tensor的数据类型,默认值为None,表示使用全局默认dtype类型。23- **dtype** (`torch.dtype`):可选参数,表示生成Tensor的数据类型,默认值为None,表示使用全局默认dtype类型。
24- **device** (`torch.device`):可选参数,表示生成Tensor的设备信息,默认值为None,表示使用当前默认device。24- **device** (`torch.device`):可选参数,表示生成Tensor的设备信息,默认值为None,表示使用当前默认device。
25 25 
26- 
27## 返回值说明26## 返回值说明
27+ 
28`Tensor`28`Tensor`
29 29 
30代表生成的特殊Tensor。30代表生成的特殊Tensor。
@@ -44,7 +44,6 @@ torch_npu.empty_with_swapped_memory(size, dtype=None, device=None) -> Tensor
44- 当安装CANN版本8.5.0及以上,且Ascend HDK版本25.5.2及以上时,该接口申请的特殊Tensor支持直接打印。44- 当安装CANN版本8.5.0及以上,且Ascend HDK版本25.5.2及以上时,该接口申请的特殊Tensor支持直接打印。
45- 当安装CANN版本小于8.5.0或者Ascend HDK版本小于25.5.2时,该接口申请的特殊Tensor不支持直接打印,此时会打印warning日志,需要查看值时要先通过`mul_`转为普通Tensor再打印。45- 当安装CANN版本小于8.5.0或者Ascend HDK版本小于25.5.2时,该接口申请的特殊Tensor不支持直接打印,此时会打印warning日志,需要查看值时要先通过`mul_`转为普通Tensor再打印。
46 46 
47- 
48## 调用示例47## 调用示例
49 48 
50单算子模式调用49单算子模式调用
Mdocs/zh/custom_APIs/torch_npu/torch_npu-erase_stream.md+3-3
@@ -1,4 +1,5 @@
1# torch_npu.erase_stream1# torch_npu.erase_stream
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -6,7 +7,6 @@
6| <term>Atlas A2 训练系列产品</term> | √ |7| <term>Atlas A2 训练系列产品</term> | √ |
7| <term>Atlas 训练系列产品</term> | √ |8| <term>Atlas 训练系列产品</term> | √ |
8 9 
9-
10## 功能说明10## 功能说明
11 11 
12- `Tensor`通过record_stream在内存池上添加的已被stream使用的标记后,可以通过该接口移除该标记。12- `Tensor`通过record_stream在内存池上添加的已被stream使用的标记后,可以通过该接口移除该标记。
@@ -17,7 +17,7 @@
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.erase_stream(tensor, stream) -> None21torch_npu.erase_stream(tensor, stream) -> None
22```22```
23 23 
@@ -27,6 +27,7 @@ torch_npu.erase_stream(tensor, stream) -> None
27- **stream** (`torch_npu.npu.Stream`):必选参数,被移除标记所属的stream。27- **stream** (`torch_npu.npu.Stream`):必选参数,被移除标记所属的stream。
28 28 
29## 返回值说明29## 返回值说明
30+ 
30`None`31`None`
31 32 
32无返回值33无返回值
@@ -37,7 +38,6 @@ torch_npu.erase_stream(tensor, stream) -> None
37 38 
38## 调用示例39## 调用示例
39 40 
40- 
41```python41```python
42>>> import torch42>>> import torch
43>>> import torch_npu43>>> import torch_npu
Mdocs/zh/custom_APIs/torch_npu/torch_npu-get_device_limit.md+3-2
@@ -1,4 +1,5 @@
1# torch.npu.get_device_limit1# torch.npu.get_device_limit
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -6,7 +7,6 @@
6|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
7|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
8 9 
9- 
10## 功能说明10## 功能说明
11 11 
12- 通过该接口,获取指定Device上的Device资源限制。12- 通过该接口,获取指定Device上的Device资源限制。
@@ -14,7 +14,7 @@
14 14 
15## 函数原型15## 函数原型
16 16 
17-```17+```python
18torch.npu.get_device_limit(device) ->Dict18torch.npu.get_device_limit(device) ->Dict
19```19```
20 20 
@@ -23,6 +23,7 @@ torch.npu.get_device_limit(device) ->Dict
23**device** (`Device`):必选参数,设置控核的卡号。23**device** (`Device`):必选参数,设置控核的卡号。
24 24 
25## 返回值说明25## 返回值说明
26+ 
26`Dict`27`Dict`
27 28 
28代表`Device`的Cube和Vector核数。29代表`Device`的Cube和Vector核数。
Mdocs/zh/custom_APIs/torch_npu/torch_npu-get_stream_limit.md+3-2
@@ -1,4 +1,5 @@
1# torch.npu.get_stream_limit1# torch.npu.get_stream_limit
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -6,7 +7,6 @@
6|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
7|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
8 9 
9- 
10## 功能说明10## 功能说明
11 11 
12- 通过该接口,获取指定Stream的Device资源限制。12- 通过该接口,获取指定Stream的Device资源限制。
@@ -15,7 +15,7 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch.npu.get_stream_limit(stream) ->Dict19torch.npu.get_stream_limit(stream) ->Dict
20```20```
21 21 
@@ -24,6 +24,7 @@ torch.npu.get_stream_limit(stream) ->Dict
24**stream** (`torch_npu.npu.Stream`):必选参数,设置控核的流。24**stream** (`torch_npu.npu.Stream`):必选参数,设置控核的流。
25 25 
26## 返回值说明26## 返回值说明
27+ 
27`Dict`28`Dict`
28 29 
29代表`stream`的Cube和Vector核数。30代表`stream`的Cube和Vector核数。
Mdocs/zh/custom_APIs/torch_npu/torch_npu-matmul_checksum.md+3-2
@@ -1,4 +1,5 @@
1# torch_npu.matmul_checksum1# torch_npu.matmul_checksum
2+ 
2## 产品支持情况3## 产品支持情况
3 4 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -12,7 +13,7 @@
12 13 
13## 函数原型14## 函数原型
14 15 
15-```16+```python
16torch_npu.matmul_checksum(a, b, c) -> Tensor17torch_npu.matmul_checksum(a, b, c) -> Tensor
17```18```
18 19 
@@ -23,6 +24,7 @@ torch_npu.matmul_checksum(a, b, c) -> Tensor
23- **c** (`Tensor`):必选输入,原生matmul计算的输出out。24- **c** (`Tensor`):必选输入,原生matmul计算的输出out。
24 25 
25## 返回值说明26## 返回值说明
27+ 
26`Tensor`28`Tensor`
27 29 
28返回NPU上的bool标量。结果为True时,标识存在aicore错误的硬件故障。30返回NPU上的bool标量。结果为True时,标识存在aicore错误的硬件故障。
@@ -33,7 +35,6 @@ torch_npu.matmul_checksum(a, b, c) -> Tensor
33 35 
34## 调用示例36## 调用示例
35 37 
36- 
37 ```python38 ```python
38 >>> import torch39 >>> import torch
39 >>> import torch_npu40 >>> import torch_npu
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_add_rms_norm.md+3-3
@@ -8,7 +8,6 @@
8| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |8| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |
9| <term>Atlas 推理系列产品</term> | √ |9| <term>Atlas 推理系列产品</term> | √ |
10 10 
11- 
12## 功能说明11## 功能说明
13 12 
14- API功能:将Add计算与RMSNorm归一化融合,常用于大模型中将残差连接后的张量进行归一化处理。13- API功能:将Add计算与RMSNorm归一化融合,常用于大模型中将残差连接后的张量进行归一化处理。
@@ -23,7 +22,8 @@
23 $$22 $$
24 23 
25## 函数原型24## 函数原型
26-```25+ 
26+```python
27torch_npu.npu_add_rms_norm(x1, x2, gamma, epsilon=1e-06) -> (Tensor, Tensor, Tensor)27torch_npu.npu_add_rms_norm(x1, x2, gamma, epsilon=1e-06) -> (Tensor, Tensor, Tensor)
28```28```
29 29 
@@ -72,4 +72,4 @@ print("y.dtype:", y.dtype)
72print("rstd:", rstd)72print("rstd:", rstd)
73print("rstd.dtype:", rstd.dtype)73print("rstd.dtype:", rstd.dtype)
74print("x:", x)74print("x:", x)
75-```75+```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_add_rms_norm_dynamic_quant.md+2-3
@@ -7,7 +7,6 @@
7| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |7| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |
8| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |8| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |
9 9 
10- 
11## 功能说明10## 功能说明
12 11 
13- API功能:RmsNorm算子是大模型常用的归一化操作,相比LayerNorm算子,其去掉了减去均值的部分。DynamicQuant算子则是为输入张量进行对称动态量化的算子。AddRmsNormDynamicQuant算子将RmsNorm前的Add算子和RmsNorm归一化输出给到的1个或2个DynamicQuant算子融合起来,减少搬入搬出操作。12- API功能:RmsNorm算子是大模型常用的归一化操作,相比LayerNorm算子,其去掉了减去均值的部分。DynamicQuant算子则是为输入张量进行对称动态量化的算子。AddRmsNormDynamicQuant算子将RmsNorm前的Add算子和RmsNorm归一化输出给到的1个或2个DynamicQuant算子融合起来,减少搬入搬出操作。
@@ -50,7 +49,6 @@
50 \end{cases}49 \end{cases}
51 $$50 $$
52 51 
53- 
54$$52$$
55 scale2Out=\begin{cases}53 scale2Out=\begin{cases}
56 row\_max(abs(input2))/127 & outputMask[1]=True\ ||\ (!outputMask\ \&\ smoothScale1Optional\ \&\ smoothScale2Optional) \\54 row\_max(abs(input2))/127 & outputMask[1]=True\ ||\ (!outputMask\ \&\ smoothScale1Optional\ \&\ smoothScale2Optional) \\
@@ -68,7 +66,8 @@ $$
68 公式中的row\_max代表每行求最大值。66 公式中的row\_max代表每行求最大值。
69 67 
70## 函数原型68## 函数原型
71-```69+ 
70+```python
72torch_npu.npu_add_rms_norm_dynamic_quant(x1, x2, gamma, *, smooth_scale1=None, smooth_scale2=None, beta=None, epsilon=1e-6, output_mask=[], y_dtype=None) -> (Tensor, Tensor, Tensor, Tensor, Tensor)71torch_npu.npu_add_rms_norm_dynamic_quant(x1, x2, gamma, *, smooth_scale1=None, smooth_scale2=None, beta=None, epsilon=1e-6, output_mask=[], y_dtype=None) -> (Tensor, Tensor, Tensor, Tensor, Tensor)
73```72```
74 73 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_add_rms_norm_quant.md+3-2
@@ -8,7 +8,6 @@
8| <term>Atlas A2 训练系列产品/Atlas 800I A2 推理产品/A200I A2 Box 异构组件</term> | √ |8| <term>Atlas A2 训练系列产品/Atlas 800I A2 推理产品/A200I A2 Box 异构组件</term> | √ |
9| <term>Atlas 推理系列产品 </term> | √ |9| <term>Atlas 推理系列产品 </term> | √ |
10 10 
11- 
12## 功能说明11## 功能说明
13 12 
14- API功能:RMSNorm是大模型常用的标准化操作,相比LayerNorm其去掉了减去均值的部分。torch_npu.npu_add_rms_norm_quant算子将RMSNorm前的Add算子以及RMSNorm后的Quantize算子融合起来,减少搬入搬出操作。13- API功能:RMSNorm是大模型常用的标准化操作,相比LayerNorm其去掉了减去均值的部分。torch_npu.npu_add_rms_norm_quant算子将RMSNorm前的Add算子以及RMSNorm后的Quantize算子融合起来,减少搬入搬出操作。
@@ -44,6 +43,7 @@
44 $$43 $$
45 44 
46## 函数原型45## 函数原型
46+ 
47```python47```python
48torch_npu.npu_add_rms_norm_quant(x1, x2, gamma, scales1, zero_points1, beta=None, scales2=None, zero_points2=None, *, axis=-1, epsilon=1e-06, div_mode=True) -> (y1, y2, x)48torch_npu.npu_add_rms_norm_quant(x1, x2, gamma, scales1, zero_points1, beta=None, scales2=None, zero_points2=None, *, axis=-1, epsilon=1e-06, div_mode=True) -> (y1, y2, x)
49```49```
@@ -69,6 +69,7 @@ torch_npu.npu_add_rms_norm_quant(x1, x2, gamma, scales1, zero_points1, beta=None
69 - **x**`Tensor`):表示`x1``x2`的和,公式中的$x$。数据格式支持$ND$,支持非连续的Tensor。数据类型和shape与输入`x1`一致。69 - **x**`Tensor`):表示`x1``x2`的和,公式中的$x$。数据格式支持$ND$,支持非连续的Tensor。数据类型和shape与输入`x1`一致。
70 70 
71## 约束说明71## 约束说明
72+ 
72- <term>Atlas 推理系列产品</term>`x1``x2`的最后一维数据个数不能小于32。`gamma``beta``scales1``zero_points1``scales2``zero_points2`的数据个数不能小于32。73- <term>Atlas 推理系列产品</term>`x1``x2`的最后一维数据个数不能小于32。`gamma``beta``scales1``zero_points1``scales2``zero_points2`的数据个数不能小于32。
73 74 
74- **边界值场景说明**75- **边界值场景说明**
@@ -141,4 +142,4 @@ def test_npu_add_rms_norm_quant():
141 142 
142if __name__ == "__main__":143if __name__ == "__main__":
143 test_npu_add_rms_norm_quant()144 test_npu_add_rms_norm_quant()
144-```145+```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_advance_step_flashattn.md+5-2
@@ -29,9 +29,10 @@
29 seqLens[i] = inputPositions[i] + 1 \\29 seqLens[i] = inputPositions[i] + 1 \\
30 slotMapping[i] = ({blockTables}[i] + blockTablesStride * i) * blockSize + (inputPositions[i]\%blockSize)30 slotMapping[i] = ({blockTables}[i] + blockTablesStride * i) * blockSize + (inputPositions[i]\%blockSize)
31 $$31 $$
32+ 
32## 函数原型33## 函数原型
33 34 
34-```35+```python
35torch_npu.npu_advance_step_flashattn(input_tokens, sampled_token_ids, input_positions, seq_lens, slot_mapping, block_tables, num_seqs, num_queries, block_size) -> ()36torch_npu.npu_advance_step_flashattn(input_tokens, sampled_token_ids, input_positions, seq_lens, slot_mapping, block_tables, num_seqs, num_queries, block_size) -> ()
36```37```
37 38 
@@ -65,6 +66,7 @@ torch_npu.npu_advance_step_flashattn(input_tokens, sampled_token_ids, input_posi
65## 调用示例66## 调用示例
66 67 
67非投机场景:68非投机场景:
69+ 
68```python70```python
69import numpy as np71import numpy as np
70 72
@@ -94,6 +96,7 @@ torch_npu.npu_advance_step_flashattn(input_tokens, sampled_token_ids, input_posi
94```96```
95 97 
96投机场景:98投机场景:
99+ 
97```python100```python
98import numpy as np101import numpy as np
99 102 
@@ -126,4 +129,4 @@ block_tables = torch.tensor(block_table, dtype=torch.int64, device="npu")
126torch_npu.npu_advance_step_flashattn(input_tokens, sampled_token_ids, input_positions,129torch_npu.npu_advance_step_flashattn(input_tokens, sampled_token_ids, input_positions,
127 seq_lens, slot_mappings, block_tables, num_seqs,130 seq_lens, slot_mappings, block_tables, num_seqs,
128 num_seqs, block_size, spec_token=spec_tokens, accepted_num=accepted_nums)131 num_seqs, block_size, spec_token=spec_tokens, accepted_num=accepted_nums)
129-```132+```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_all_gather_base_mm.md+25-26
@@ -1,7 +1,5 @@
1# torch\_npu.npu\_all\_gather\_base\_mm1# torch\_npu.npu\_all\_gather\_base\_mm
2 2 
3- 
4- 
5## 产品支持情况3## 产品支持情况
6 4 
7| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -11,9 +9,9 @@
11 9 
12## 功能说明<a name="zh-cn_topic_0000001694916914_section14441124184110"></a>10## 功能说明<a name="zh-cn_topic_0000001694916914_section14441124184110"></a>
13 11 
14-- API功能:TP切分场景下,融合`allgather`和`matmul`,实现通信和计算流水并行。12+- API功能:TP切分场景下,融合`allgather`和`matmul`,实现通信和计算流水并行。
15 13 
16-- 计算公式:14+- 计算公式:
17 $x1$代表输入`x1`15 $x1$代表输入`x1`
18 16
19 基础场景:17 基础场景:
@@ -36,44 +34,46 @@
36 34 
37## 函数原型35## 函数原型
38 36 
39-```37+```python
40torch_npu.npu_all_gather_base_mm(x1, x2, hcom, world_size, bias=None, x1_scale=None, x2_scale=None, gather_index=0, gather_output=True, comm_turn=0, output_dtype=None, comm_mode=None) -> tuple[Tensor, Tensor]38torch_npu.npu_all_gather_base_mm(x1, x2, hcom, world_size, bias=None, x1_scale=None, x2_scale=None, gather_index=0, gather_output=True, comm_turn=0, output_dtype=None, comm_mode=None) -> tuple[Tensor, Tensor]
41```39```
42 40 
43## 参数说明41## 参数说明
44 42 
45-- **x1** (`Tensor`):必选参数,表示矩阵乘法中的左矩阵,数据类型支持`float16`、`bfloat16`、`int8`,数据格式支持ND,输入shape支持2维,形如\(m, k\),轴满足matmul算子入参要求,第二轴与`x2`的第一轴相等,且k的取值范围为\[256, 65535\)。43+- **x1** (`Tensor`):必选参数,表示矩阵乘法中的左矩阵,数据类型支持`float16`、`bfloat16`、`int8`,数据格式支持ND,输入shape支持2维,形如\(m, k\),轴满足matmul算子入参要求,第二轴与`x2`的第一轴相等,且k的取值范围为\[256, 65535\)。
46-- **x2** (`Tensor`):必选参数,表示矩阵乘法中的右矩阵,数据类型需要和`x1`保持一致,数据格式支持$ND$、$NZ$。$NZ$仅在`comm_mode`为`aiv`时支持。输入shape支持2维,形如\(k, n\),轴满足matmul算子入参要求,第一轴与`x1`的第二轴相等,且k的取值范围为\[256, 65535\)。44+- **x2** (`Tensor`):必选参数,表示矩阵乘法中的右矩阵,数据类型需要和`x1`保持一致,数据格式支持$ND$、$NZ$。$NZ$仅在`comm_mode`为`aiv`时支持。输入shape支持2维,形如\(k, n\),轴满足matmul算子入参要求,第一轴与`x1`的第二轴相等,且k的取值范围为\[256, 65535\)。
47-- **hcom** (`string`):必选参数,通信域handle名,通过get\_hccl\_comm\_name接口获取。45+- **hcom** (`string`):必选参数,通信域handle名,通过get\_hccl\_comm\_name接口获取。
48-- **world\_size** (`int`):必选参数,通信域内的rank总数。46+- **world\_size** (`int`):必选参数,通信域内的rank总数。
49- - <term>Atlas A2 训练系列产品</term>:支持2、4、8卡,支持hccs链路all mesh组网(每张卡和其它卡两两相连)。47+ - <term>Atlas A2 训练系列产品</term>:支持2、4、8卡,支持hccs链路all mesh组网(每张卡和其它卡两两相连)。
50- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持2、4、8、16、32卡,支持hccs链路double ring组网(多张卡按顺序组成一个圈,每张卡只和左右卡相连)。48+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持2、4、8、16、32卡,支持hccs链路double ring组网(多张卡按顺序组成一个圈,每张卡只和左右卡相连)。
51-- **bias** (`Tensor`):可选参数,数据类型支持float16、bfloat16,数据格式支持ND格式。数据类型需要和`x1`保持一致。bias仅支持一维,且维度大小与`output`的第1维大小相同。**当前版本暂不支持bias输入为非0的场景。**49+- **bias** (`Tensor`):可选参数,数据类型支持float16、bfloat16,数据格式支持ND格式。数据类型需要和`x1`保持一致。bias仅支持一维,且维度大小与`output`的第1维大小相同。**当前版本暂不支持bias输入为非0的场景。**
52- **x1\_scale** (`Tensor`):可选参数,mm左矩阵反量化参数。数据类型支持`float32`,数据格式支持$ND$格式。数据维度为\(m, 1\),支持pertoken量化。50- **x1\_scale** (`Tensor`):可选参数,mm左矩阵反量化参数。数据类型支持`float32`,数据格式支持$ND$格式。数据维度为\(m, 1\),支持pertoken量化。
53- **x2\_scale** (`Tensor`):可选参数,mm右矩阵反量化参数。数据类型支持`float32`、`int64`,数据格式支持$ND$格式。数据维度为\(1, n\),支持perchannel量化。如需传入`int64`数据类型的,需要提前调用torch_npu.npu_trans_quant_param来获取`int64`数据类型的`x2_scale`。51- **x2\_scale** (`Tensor`):可选参数,mm右矩阵反量化参数。数据类型支持`float32`、`int64`,数据格式支持$ND$格式。数据维度为\(1, n\),支持perchannel量化。如需传入`int64`数据类型的,需要提前调用torch_npu.npu_trans_quant_param来获取`int64`数据类型的`x2_scale`。
54-- **gather\_index** (`int`):可选参数,表示gather操作对象,0表示对`x1`做gather,1表示对`x2`做gather。默认值0。**当前版本仅支持输入0。**52+- **gather\_index** (`int`):可选参数,表示gather操作对象,0表示对`x1`做gather,1表示对`x2`做gather。默认值0。**当前版本仅支持输入0。**
55-- **gather\_output** (`bool`):可选参数,表示是否需要gather输出。默认值True。53+- **gather\_output** (`bool`):可选参数,表示是否需要gather输出。默认值True。
56-- **comm\_turn** (`int`):可选参数,表示rank间通信切分粒度,默认值为0,表示默认的切分方式。**当前版本仅支持输入0。**54+- **comm\_turn** (`int`):可选参数,表示rank间通信切分粒度,默认值为0,表示默认的切分方式。**当前版本仅支持输入0。**
57- **output_dtype** (`ScalarType`):可选参数,表示第一个输出的数据类型。仅支持在量化场景且`x1_scale`和`x2_scale`均为`float32`时,可指定输出数据类型为`bfloat16`或`float16`,默认值为`bfloat16`。55- **output_dtype** (`ScalarType`):可选参数,表示第一个输出的数据类型。仅支持在量化场景且`x1_scale`和`x2_scale`均为`float32`时,可指定输出数据类型为`bfloat16`或`float16`,默认值为`bfloat16`。
58- **comm\_mode** (`string`):可选参数,表示通信模式,支持`ai_cpu`、`aiv`两种模式。`ai_cpu`模式仅支持基础场景。`aiv`模式支持基础场景和量化场景。默认值为`ai_cpu`。56- **comm\_mode** (`string`):可选参数,表示通信模式,支持`ai_cpu`、`aiv`两种模式。`ai_cpu`模式仅支持基础场景。`aiv`模式支持基础场景和量化场景。默认值为`ai_cpu`。
59 57 
60## 返回值说明<a name="zh-cn_topic_0000001694916914_section15236153161410"></a>58## 返回值说明<a name="zh-cn_topic_0000001694916914_section15236153161410"></a>
61-- **output** (`Tensor`):第一个输出Tensor,allgather+matmul的结果。59+ 
60+- **output** (`Tensor`):第一个输出Tensor,allgather+matmul的结果。
62基础场景时数据类型和`x1`保持一致。61基础场景时数据类型和`x1`保持一致。
63量化场景下,`x2_scale``int64`数据类型时,输出数据类型为`float16``x1_scale``x2_scale`均为`float32`时,输出数据类型由`output_dtype`指定,默认为`bfloat16`62量化场景下,`x2_scale``int64`数据类型时,输出数据类型为`float16``x1_scale``x2_scale`均为`float32`时,输出数据类型由`output_dtype`指定,默认为`bfloat16`
64-- **gather_out** (`Tensor`):第二个输出Tensor,allgather的结果,由`gather_output`参数控制是否输出,`gather_output`为False时,返回空Tensor。63+- **gather_out** (`Tensor`):第二个输出Tensor,allgather的结果,由`gather_output`参数控制是否输出,`gather_output`为False时,返回空Tensor。
65 64 
66## 约束说明65## 约束说明
67-- `x1`不支持输入转置后的tensor,`x2`转置后输入,需要满足shape的第一维大小与`x1`的最后一维相同,满足matmul的计算条件。66+ 
68-- `comm_mode``ai_cpu`时:67+- `x1`不支持输入转置后的tensor,`x2`转置后输入,需要满足shape的第一维大小与`x1`的最后一维相同,满足matmul的计算条件。
69- - 该接口支持训练场景下使用。68+- `comm_mode`为`ai_cpu`时:
70- - 该接口支持图模式 69+ - 该接口支持训练场景下使用
71- - <term>Atlas A2 训练系列产品</term>:一个模型中的通算融合算子(AllGatherMatmul、MatmulReduceScatter、MatmulAllReduce),仅支持相同通信域70+ - 该接口支持图模式
72-- `comm_mode`为`aiv`时,训练和推理场景均可使用71+ - <term>Atlas A2 训练系列产品</term>:一个模型中的通算融合算子(AllGatherMatmul、MatmulReduceScatter、MatmulAllReduce),仅支持相同通信域
72+- `comm_mode``aiv`时,训练和推理场景均可使用。
73 73 
74## 调用示例<a name="zh-cn_topic_0000001694916914_section14459801435"></a>74## 调用示例<a name="zh-cn_topic_0000001694916914_section14459801435"></a>
75 75 
76-- 单算子模式调用76+- 单算子模式调用
77 77 
78 ```python78 ```python
79 import torch79 import torch
@@ -109,7 +109,7 @@ torch_npu.npu_all_gather_base_mm(x1, x2, hcom, world_size, bias=None, x1_scale=N
109 mp.spawn(run_all_gather_base_mm, args=(worksize, master_ip, master_port, x1_shape, x2_shape, dtype), nprocs=worksize)109 mp.spawn(run_all_gather_base_mm, args=(worksize, master_ip, master_port, x1_shape, x2_shape, dtype), nprocs=worksize)
110 ```110 ```
111 111 
112-- 图模式调用112+- 图模式调用
113 113 
114 ```python114 ```python
115 import torch115 import torch
@@ -165,4 +165,3 @@ torch_npu.npu_all_gather_base_mm(x1, x2, hcom, world_size, bias=None, x1_scale=N
165 dtype = torch.float16165 dtype = torch.float16
166 mp.spawn(run_all_gather_base_mm, args=(worksize, master_ip, master_port, x1_shape, x2_shape, dtype), nprocs=worksize)166 mp.spawn(run_all_gather_base_mm, args=(worksize, master_ip, master_port, x1_shape, x2_shape, dtype), nprocs=worksize)
167 ```167 ```
168- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_all_to_all_matmul.md+21-21
@@ -20,40 +20,40 @@
20 20 
21## 函数原型21## 函数原型
22 22 
23-```23+```python
24torch_npu.npu_all_to_all_matmul(x1, x2, hcom, world_size, bias=None, all2all_axes=None, all2all_out_flag=True) -> (Tensor, Tensor)24torch_npu.npu_all_to_all_matmul(x1, x2, hcom, world_size, bias=None, all2all_axes=None, all2all_out_flag=True) -> (Tensor, Tensor)
25```25```
26 26 
27## 参数说明27## 参数说明
28 28 
29-- **x1**(`Tensor`):必选输入,表示融合算子的左矩阵输入,也是Matmul计算的左矩阵,对应公式中的x1。数据类型支持bfloat16、float16,维度只能为2D,shape为(BS, H),数据格式支持ND,不支持非连续Tensor,支持第一维度为0的空Tensor。29+- **x1**(`Tensor`):必选输入,表示融合算子的左矩阵输入,也是Matmul计算的左矩阵,对应公式中的x1。数据类型支持bfloat16、float16,维度只能为2D,shape为(BS, H),数据格式支持ND,不支持非连续Tensor,支持第一维度为0的空Tensor。
30-- **x2**(`Tensor`):必选输入,表示融合算子的右矩阵输入,也是Matmul计算的右矩阵,对应公式中的x2。数据类型与x1一致,维度只能为2D,shape为(H*rankSize, N),数据格式支持ND,支持转置非连续Tensor。30+- **x2**(`Tensor`):必选输入,表示融合算子的右矩阵输入,也是Matmul计算的右矩阵,对应公式中的x2。数据类型与x1一致,维度只能为2D,shape为(H*rankSize, N),数据格式支持ND,支持转置非连续Tensor。
31-- **hcom**(`str`):必选输入,Host侧标识列组的字符串,即通信域名称,通过get_hccl_comm_name接口获取。31+- **hcom**(`str`):必选输入,Host侧标识列组的字符串,即通信域名称,通过get_hccl_comm_name接口获取。
32-- **world_size**(`int`):必选输入,通信域内的rank总数,对应公式中的rankSize,支持范围[2, 4, 8, 16]。32+- **world_size**(`int`):必选输入,通信域内的rank总数,对应公式中的rankSize,支持范围[2, 4, 8, 16]。
33-- **bias**(`Tensor`):可选输入,矩阵乘运算后累加的偏置,对应公式中的bias。数据类型由输入x1和x2决定,当x1和x2为float16时,bias的数据类型为float16;当x1和x2为bfloat16时,bias的数据类型为float32。维度只能为1D,shape为(N),数据类型支持ND。33+- **bias**(`Tensor`):可选输入,矩阵乘运算后累加的偏置,对应公式中的bias。数据类型由输入x1和x2决定,当x1和x2为float16时,bias的数据类型为float16;当x1和x2为bfloat16时,bias的数据类型为float32。维度只能为1D,shape为(N),数据类型支持ND。
34-- **all2all_axes**(`List[int]`):可选输入,AlltoAll和Permute数据交换的方向,支持为空或者[-2, -1],表示将Matmul结果由(BS, H)转为(BS/rankSize, H*rankSize)。34+- **all2all_axes**(`List[int]`):可选输入,AlltoAll和Permute数据交换的方向,支持为空或者[-2, -1],表示将Matmul结果由(BS, H)转为(BS/rankSize, H*rankSize)。
35-- **all2all_out_flag**(`bool`):可选输入,表示是否输出AlltoAll和Permute后的结果,默认为True。35+- **all2all_out_flag**(`bool`):可选输入,表示是否输出AlltoAll和Permute后的结果,默认为True。
36 36 
37## 返回值说明37## 返回值说明
38 38 
39-- **y**(`Tensor`):计算输出,表示最终的计算结果,数据类型与输入x1或者x2保持一致,支持2维,shape为(BS/rankSize, N),数据格式支持ND,不支持非连续的Tensor。39+- **y**(`Tensor`):计算输出,表示最终的计算结果,数据类型与输入x1或者x2保持一致,支持2维,shape为(BS/rankSize, N),数据格式支持ND,不支持非连续的Tensor。
40-- **all2all_out**(`Tensor`):计算输出,表示AlltoAll和Permute后的结果,公式中的permutedOut。当all2all_out_flag为True时输出实际Tensor,数据类型与输入x1或者x2保持一致,支持2维,shape为(BS/rankSize, H*rankSize),数据格式支持ND,不支持非连续的Tensor。40+- **all2all_out**(`Tensor`):计算输出,表示AlltoAll和Permute后的结果,公式中的permutedOut。当all2all_out_flag为True时输出实际Tensor,数据类型与输入x1或者x2保持一致,支持2维,shape为(BS/rankSize, H*rankSize),数据格式支持ND,不支持非连续的Tensor。
41 41 
42## 约束说明42## 约束说明
43 43 
44-- 该接口支持训练、推理场景下使用。44+- 该接口支持训练、推理场景下使用。
45-- A3场景下,该接口支持单算子模式,不支持图模式。45+- A3场景下,该接口支持单算子模式,不支持图模式。
46-- 除x1以外的输入参数均不支持空Tensor。46+- 除x1以外的输入参数均不支持空Tensor。
47-- 通信域名称hcom不支持传入空字符串,长度取值范围为[1, 127]。47+- 通信域名称hcom不支持传入空字符串,长度取值范围为[1, 127]。
48-- 输入参数Tensor中shape使用的变量说明:48+- 输入参数Tensor中shape使用的变量说明:
49- - BS:输入左矩阵的第一维度大小,表示输入序列sequence的条数,取值范围为[0, 2147483647],必须整除rankSize。49+ - BS:输入左矩阵的第一维度大小,表示输入序列sequence的条数,取值范围为[0, 2147483647],必须整除rankSize。
50- - H:输入左矩阵的第二维度大小,表示隐藏层维度,取值范围受NPU卡数限制,H*rankSize取值范围为[2, 65535]。50+ - H:输入左矩阵的第二维度大小,表示隐藏层维度,取值范围受NPU卡数限制,H*rankSize取值范围为[2, 65535]。
51- - H*rankSize:输入右矩阵的第一维度大小,表示输入左矩阵经过AlltoAll通信后隐藏层维度大小,取值范围为[2, 65535]。51+ - H*rankSize:输入右矩阵的第一维度大小,表示输入左矩阵经过AlltoAll通信后隐藏层维度大小,取值范围为[2, 65535]。
52- - N:输入右矩阵的第二维度大小,表示输出序列sequence的长度,取值范围为[1, 2147483647]。52+ - N:输入右矩阵的第二维度大小,表示输出序列sequence的长度,取值范围为[1, 2147483647]。
53 53 
54## 调用示例54## 调用示例
55 55 
56-- 单算子模式调用56+- 单算子模式调用
57 57 
58 ```python58 ```python
59 import torch59 import torch
@@ -93,4 +93,4 @@ torch_npu.npu_all_to_all_matmul(x1, x2, hcom, world_size, bias=None, all2all_axe
93 args=(worksize, master_ip, master_port, x1_shape, x2_shape),93 args=(worksize, master_ip, master_port, x1_shape, x2_shape),
94 nprocs=worksize,94 nprocs=worksize,
95 )95 )
96- ```96+ ```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_alltoallv_gmm.md+42-43
@@ -10,75 +10,75 @@
10 10 
11- API功能:MoE(Mixture of Experts,混合专家模型)网络中,完成路由专家AlltoAllv、Permute、GroupedMatMul融合并实现与共享专家MatMul并行融合,先通信后计算。11- API功能:MoE(Mixture of Experts,混合专家模型)网络中,完成路由专家AlltoAllv、Permute、GroupedMatMul融合并实现与共享专家MatMul并行融合,先通信后计算。
12 12 
13-- 路由专家计算公式:13+- 路由专家计算公式:
14 14 
15 ![](../../figures/zh-cn_formulaimage_0000002357766385.png)15 ![](../../figures/zh-cn_formulaimage_0000002357766385.png)
16 16 
17- - ata\_out是gmm\_x进行AlltoAllv通信的输出结果,后续用于Permute计算。17+ - ata\_out是gmm\_x进行AlltoAllv通信的输出结果,后续用于Permute计算。
18- - permute\_out是ata\_out进行Permute计算的输出结果,作为路由专家进行GroupedMatMul计算的左矩阵。18+ - permute\_out是ata\_out进行Permute计算的输出结果,作为路由专家进行GroupedMatMul计算的左矩阵。
19- - gmm\_weight指路由专家进行GroupedMatMul计算的右矩阵。19+ - gmm\_weight指路由专家进行GroupedMatMul计算的右矩阵。
20- - gmm\_y指路由专家进行GroupedMatMul计算的输出。20+ - gmm\_y指路由专家进行GroupedMatMul计算的输出。
21 21 
22-- 共享专家计算公式:22+- 共享专家计算公式:
23 23 
24 ![](../../figures/zh-cn_formulaimage_0000002323547858.png)24 ![](../../figures/zh-cn_formulaimage_0000002323547858.png)
25 25 
26- - mm\_x指共享专家MatMul计算的左矩阵。26+ - mm\_x指共享专家MatMul计算的左矩阵。
27- - mm\_weight指共享专家MatMul计算的右矩阵。27+ - mm\_weight指共享专家MatMul计算的右矩阵。
28- - mm\_y指共享专家MatMul计算的输出。28+ - mm\_y指共享专家MatMul计算的输出。
29 29 
30## 函数原型<a name="zh-cn_topic_0000002282815538_section45077510411"></a>30## 函数原型<a name="zh-cn_topic_0000002282815538_section45077510411"></a>
31 31 
32-```32+```python
33torch_npu.npu_alltoallv_gmm(gmm_x, gmm_weight, hcom, ep_world_size, send_counts, recv_counts, *, send_counts_tensor=None, recv_counts_tensor=None, mm_x=None, mm_weight=None, trans_gmm_weight=False, trans_mm_weight=False, permute_out_flag=False) -> (Tensor, Tensor, Tensor)33torch_npu.npu_alltoallv_gmm(gmm_x, gmm_weight, hcom, ep_world_size, send_counts, recv_counts, *, send_counts_tensor=None, recv_counts_tensor=None, mm_x=None, mm_weight=None, trans_gmm_weight=False, trans_mm_weight=False, permute_out_flag=False) -> (Tensor, Tensor, Tensor)
34```34```
35 35 
36## 参数说明<a name="zh-cn_topic_0000002282815538_section112637109429"></a>36## 参数说明<a name="zh-cn_topic_0000002282815538_section112637109429"></a>
37 37 
38-- **gmm\_x**(`Tensor`):必选参数,AlltoAllv通信与Permute操作后结果作为GroupedMatMul计算的左矩阵。数据类型支持`float16`、`bfloat16`,支持2维,shape为$(BSK, H1)$,数据格式支持ND。38+- **gmm\_x**(`Tensor`):必选参数,AlltoAllv通信与Permute操作后结果作为GroupedMatMul计算的左矩阵。数据类型支持`float16`、`bfloat16`,支持2维,shape为$(BSK, H1)$,数据格式支持ND。
39-- **gmm\_weight**(`Tensor`):必选参数,GroupedMatMul计算的右矩阵。数据类型与`gmm_x`保持一致,支持3维,shape为$(e, H1, N1)$,数据格式支持ND。39+- **gmm\_weight**(`Tensor`):必选参数,GroupedMatMul计算的右矩阵。数据类型与`gmm_x`保持一致,支持3维,shape为$(e, H1, N1)$,数据格式支持ND。
40-- **hcom**(`str`):必选参数,专家并行的通信域名,字符串长度要求\(0, 128\)。40+- **hcom**(`str`):必选参数,专家并行的通信域名,字符串长度要求\(0, 128\)。
41-- **ep\_world\_size**(`int`):必选参数,EP通信域size,取值支持8、16、32、64、128。41+- **ep\_world\_size**(`int`):必选参数,EP通信域size,取值支持8、16、32、64、128。
42-- **send\_counts**(`List[int]`):必选参数,表示发送给其他卡的token数,数据类型支持int,取值大小为e\*`ep_world_size`,最大为256。42+- **send\_counts**(`List[int]`):必选参数,表示发送给其他卡的token数,数据类型支持int,取值大小为e\*`ep_world_size`,最大为256。
43-- **recv\_counts**(`List[int]`):必选参数,表示接收其他卡的token数,数据类型支持int,取值大小为e\*`ep_world_size`,最大为256。43+- **recv\_counts**(`List[int]`):必选参数,表示接收其他卡的token数,数据类型支持int,取值大小为e\*`ep_world_size`,最大为256。
44-- **send\_counts\_tensor**(`Tensor`):可选参数,数据类型支持int,shape为$(e*ep\_world\_size,)$,数据格式支持ND。**当前版本暂不支持**,使用默认值即可。44+- **send\_counts\_tensor**(`Tensor`):可选参数,数据类型支持int,shape为$(e*ep\_world\_size,)$,数据格式支持ND。**当前版本暂不支持**,使用默认值即可。
45-- **recv\_counts\_tensor**(`Tensor`):可选参数,数据类型支持int,shape为$(e*ep\_world\_size,)$,数据格式支持ND。**当前版本暂不支持**,使用默认值即可。45+- **recv\_counts\_tensor**(`Tensor`):可选参数,数据类型支持int,shape为$(e*ep\_world\_size,)$,数据格式支持ND。**当前版本暂不支持**,使用默认值即可。
46-- **mm\_x**(`Tensor`):可选参数,共享专家MatMul计算中的左矩阵。当需要融合共享专家矩阵计算时,该参数必选,数据类型支持`float16`、`bfloat16`,支持2维,shape为$(BS, H2)$。46+- **mm\_x**(`Tensor`):可选参数,共享专家MatMul计算中的左矩阵。当需要融合共享专家矩阵计算时,该参数必选,数据类型支持`float16`、`bfloat16`,支持2维,shape为$(BS, H2)$。
47-- **mm\_weight**(`Tensor`):可选参数,共享专家MatMul计算中的右矩阵。当需要融合共享专家矩阵计算时,该参数必选,数据类型与`mm_x`保持一致,支持2维,shape为$(H2, N2)$。47+- **mm\_weight**(`Tensor`):可选参数,共享专家MatMul计算中的右矩阵。当需要融合共享专家矩阵计算时,该参数必选,数据类型与`mm_x`保持一致,支持2维,shape为$(H2, N2)$。
48-- **trans\_gmm\_weight**(`bool`):可选参数,GroupedMatMul的右矩阵是否需要转置,true表示需要转置,false表示不转置。48+- **trans\_gmm\_weight**(`bool`):可选参数,GroupedMatMul的右矩阵是否需要转置,true表示需要转置,false表示不转置。
49-- **trans\_mm\_weight**(`bool`):可选参数,共享专家MatMul的右矩阵是否需要转置,true表示需要转置,false表示不转置。49+- **trans\_mm\_weight**(`bool`):可选参数,共享专家MatMul的右矩阵是否需要转置,true表示需要转置,false表示不转置。
50-- **permute\_out\_flag**(`bool`):可选参数,Permute结果是否需要输出,true表明需要输出,false表明不需要输出。50+- **permute\_out\_flag**(`bool`):可选参数,Permute结果是否需要输出,true表明需要输出,false表明不需要输出。
51 51 
52## 返回值说明<a name="zh-cn_topic_0000002282815538_section22231435517"></a>52## 返回值说明<a name="zh-cn_topic_0000002282815538_section22231435517"></a>
53 53 
54-- **gmm\_y**(`Tensor`):计算输出,表示最终的计算结果,数据类型与输入`gmm_x`保持一致,支持2维,shape为$(A, N1)$。54+- **gmm\_y**(`Tensor`):计算输出,表示最终的计算结果,数据类型与输入`gmm_x`保持一致,支持2维,shape为$(A, N1)$。
55-- **mm\_y**(`Tensor`):计算输出,共享专家MatMul的输出,数据类型与`mm_x`保持一致,支持2维,shape为$(BS, N2)$。仅当传入`mm_x`与`mm_weight`才输出。55+- **mm\_y**(`Tensor`):计算输出,共享专家MatMul的输出,数据类型与`mm_x`保持一致,支持2维,shape为$(BS, N2)$。仅当传入`mm_x`与`mm_weight`才输出。
56-- **permute\_out**(`Tensor`):计算输出,Permute之后的输出,数据类型与`gmm_x`保持一致。56+- **permute\_out**(`Tensor`):计算输出,Permute之后的输出,数据类型与`gmm_x`保持一致。
57 57 
58## 约束说明<a name="zh-cn_topic_0000002282815538_section12345537164214"></a>58## 约束说明<a name="zh-cn_topic_0000002282815538_section12345537164214"></a>
59 59 
60-- 该接口支持推理场景下使用。60+- 该接口支持推理场景下使用。
61-- 该接口支持图模式。61+- 该接口支持图模式。
62-- 单卡通信量取值大于等于2MB。62+- 单卡通信量取值大于等于2MB。
63-- 输入参数Tensor中shape使用的变量说明:63+- 输入参数Tensor中shape使用的变量说明:
64- - BSK:本卡发送的token数(BS\*K=BSK),是send\_counts参数累加之和,取值范围\(0, 52428800\)。64+ - BSK:本卡发送的token数(BS\*K=BSK),是send\_counts参数累加之和,取值范围\(0, 52428800\)。
65- - H1:表示路由专家hidden size隐藏层大小,取值范围\(0, 65536\)。65+ - H1:表示路由专家hidden size隐藏层大小,取值范围\(0, 65536\)。
66 66 
67- - H2:表示共享专家hidden size隐藏层大小,取值范围\(0, 12288\]。67+ - H2:表示共享专家hidden size隐藏层大小,取值范围\(0, 12288\]。
68- - e:表示单卡上专家个数,e<=32,e\*ep\_world\_size最大支持256。68+ - e:表示单卡上专家个数,e<=32,e\*ep\_world\_size最大支持256。
69 69 
70- - N1:表示路由专家的head\_num,取值范围\(0, 65536\)。70+ - N1:表示路由专家的head\_num,取值范围\(0, 65536\)。
71- - N2:表示共享专家的head\_num,取值范围\(0, 65536\)。71+ - N2:表示共享专家的head\_num,取值范围\(0, 65536\)。
72 72 
73- - BS:表示batch sequence size。73+ - BS:表示batch sequence size。
74- - K:表示选取topK个专家,K的范围\[2, 8\]。74+ - K:表示选取topK个专家,K的范围\[2, 8\]。
75 75 
76- - A:本卡收到的token数,是recv\_counts参数累加之和。76+ - A:本卡收到的token数,是recv\_counts参数累加之和。
77- - EP通信域内所有卡上的A累加和等于所有卡上的BSK累加和。77+ - EP通信域内所有卡上的A累加和等于所有卡上的BSK累加和。
78 78 
79## 调用示例<a name="zh-cn_topic_0000002282815538_section14459801435"></a>79## 调用示例<a name="zh-cn_topic_0000002282815538_section14459801435"></a>
80 80 
81-- 单算子模式调用81+- 单算子模式调用
82 82 
83 ```python83 ```python
84 import torch84 import torch
@@ -129,7 +129,7 @@ torch_npu.npu_alltoallv_gmm(gmm_x, gmm_weight, hcom, ep_world_size, send_counts,
129 mp.spawn(run_npu_alltoallv_gmm, args=(epWorkSize, master_ip, master_port, gmm_x_shape, gmm_weight_shape, send_counts, recv_counts, dtype), nprocs=epWorkSize)129 mp.spawn(run_npu_alltoallv_gmm, args=(epWorkSize, master_ip, master_port, gmm_x_shape, gmm_weight_shape, send_counts, recv_counts, dtype), nprocs=epWorkSize)
130 ```130 ```
131 131 
132-- 图模式调用132+- 图模式调用
133 133 
134 ```python134 ```python
135 import torch135 import torch
@@ -206,4 +206,3 @@ torch_npu.npu_alltoallv_gmm(gmm_x, gmm_weight, hcom, ep_world_size, send_counts,
206 dtype = torch.float16206 dtype = torch.float16
207 mp.spawn(run_npu_alltoallv_gmm, args=(epWorkSize, master_ip, master_port, gmm_x_shape, gmm_weight_shape, send_counts, recv_counts, dtype), nprocs=epWorkSize)207 mp.spawn(run_npu_alltoallv_gmm, args=(epWorkSize, master_ip, master_port, gmm_x_shape, gmm_weight_shape, send_counts, recv_counts, dtype), nprocs=epWorkSize)
208 ```208 ```
209- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_anti_quant.md+2-1
@@ -21,7 +21,7 @@
21 21 
22## 函数原型22## 函数原型
23 23 
24-```24+```python
25torch_npu.npu_anti_quant(x, scale, *, offset=None, dst_dtype=None, src_dtype=None) -> Tensor25torch_npu.npu_anti_quant(x, scale, *, offset=None, dst_dtype=None, src_dtype=None) -> Tensor
26```26```
27 27 
@@ -52,6 +52,7 @@ torch_npu.npu_anti_quant(x, scale, *, offset=None, dst_dtype=None, src_dtype=Non
52 - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`quint4x2``int8`52 - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`quint4x2``int8`
53 53 
54## 返回值说明54## 返回值说明
55+ 
55`Tensor`56`Tensor`
56 57 
57代表`npu_anti_quant`的计算结果,对应公式中的*out*。支持非连续的Tensor,支持空Tensor。58代表`npu_anti_quant`的计算结果,对应公式中的*out*。支持非连续的Tensor,支持空Tensor。
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_attention_to_ffn.md+33-34
@@ -12,60 +12,59 @@
12 12 
13## 函数原型13## 函数原型
14 14 
15-```15+```python
16torch_npu.npu_attention_to_ffn(x, session_id, micro_batch_id, layer_id, expert_ids, expert_rank_table, group, world_size, ffn_token_info_table_shape, ffn_token_data_shape, attn_token_info_table_shape, moe_expert_num, *, scales=None, active_mask=None, quant_mode=0, sync_flag=0, ffn_start_rank_id=0) -> ()16torch_npu.npu_attention_to_ffn(x, session_id, micro_batch_id, layer_id, expert_ids, expert_rank_table, group, world_size, ffn_token_info_table_shape, ffn_token_data_shape, attn_token_info_table_shape, moe_expert_num, *, scales=None, active_mask=None, quant_mode=0, sync_flag=0, ffn_start_rank_id=0) -> ()
17```17```
18 18 
19## 参数说明19## 参数说明
20 20 
21-- **x** (`Tensor`):必选参数,表示计算使用的token数据,需根据`expert_ids`和`expert_rank_table`来发送给其他卡。要求为3维张量,shape为\(X, BS, H\),表示有X个microBatch,每个microBatch里有BS个token,数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。21+- **x** (`Tensor`):必选参数,表示计算使用的token数据,需根据`expert_ids`和`expert_rank_table`来发送给其他卡。要求为3维张量,shape为\(X, BS, H\),表示有X个microBatch,每个microBatch里有BS个token,数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。
22-- **session\_id** (`Tensor`):必选参数,Attention域本卡ID,要求为1维张量,shape为\(X, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。22+- **session\_id** (`Tensor`):必选参数,Attention域本卡ID,要求为1维张量,shape为\(X, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。
23-- **micro\_batch\_id** (`Tensor`):必选参数,当前microBatch组的ID,,要求为1维张量,shape为\(X, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。23+- **micro\_batch\_id** (`Tensor`):必选参数,当前microBatch组的ID,,要求为1维张量,shape为\(X, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。
24-- **layer\_id** (`Tensor`):必选参数,模型层数ID,要求为1维张量,shape为\(X, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。24+- **layer\_id** (`Tensor`):必选参数,模型层数ID,要求为1维张量,shape为\(X, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。
25-- **expert\_ids** (`Tensor`):必选参数,每个micro batch组中每个token的topK个专家索引,决定每个token要发给哪些专家。要求为3维张量,shape为\(X, BS, K\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。25+- **expert\_ids** (`Tensor`):必选参数,每个micro batch组中每个token的topK个专家索引,决定每个token要发给哪些专家。要求为3维张量,shape为\(X, BS, K\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。
26-- **expert\_rank\_table** (`Tensor`):必选参数,每个micro batch组中专家Id到FFN卡专家部署的映射表,外部需保证值正确。要求为3维张量,shape为\(L, shared\_expert\_num + moe\_expert\_num, M\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。26+- **expert\_rank\_table** (`Tensor`):必选参数,每个micro batch组中专家Id到FFN卡专家部署的映射表,外部需保证值正确。要求为3维张量,shape为\(L, shared\_expert\_num + moe\_expert\_num, M\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。
27-- **group** (`str`):必选参数,通信域名称,专家并行的通信域。字符串长度范围为\[1,128\)。27+- **group** (`str`):必选参数,通信域名称,专家并行的通信域。字符串长度范围为\[1,128\)。
28-- **world\_size**(`int`):必选参数,通信域size。取值支持\[2, 768\]。28+- **world\_size**(`int`):必选参数,通信域size。取值支持\[2, 768\]。
29-- **ffn\_token\_info\_table\_shape** (`List(int)`):必选参数,表示FFN卡上token信息表格shape大小,长度为3,包括Attention节点的数量、microBatchSize的大小以及每个token对应的相关发送状态信息shape的大小。29+- **ffn\_token\_info\_table\_shape** (`List(int)`):必选参数,表示FFN卡上token信息表格shape大小,长度为3,包括Attention节点的数量、microBatchSize的大小以及每个token对应的相关发送状态信息shape的大小。
30-- **ffn\_token\_data\_shape** (`List(int)`):必选参数,表示FFN卡上token数据表格shape大小,长度为5,包括Attention节点的数量、microBatchSize的大小、batchSize大小、每个token需发送的专家数量(包括共享专家)、单个token的长度。30+- **ffn\_token\_data\_shape** (`List(int)`):必选参数,表示FFN卡上token数据表格shape大小,长度为5,包括Attention节点的数量、microBatchSize的大小、batchSize大小、每个token需发送的专家数量(包括共享专家)、单个token的长度。
31-- **attn\_token\_info\_table\_shape** (`List(int)`):必选参数,表示Attention卡上token信息表格shape大小,长度为3,包括microBatchSize的大小、batchSize大小、每个token需发送的专家数量(包括共享专家)。31+- **attn\_token\_info\_table\_shape** (`List(int)`):必选参数,表示Attention卡上token信息表格shape大小,长度为3,包括microBatchSize的大小、batchSize大小、每个token需发送的专家数量(包括共享专家)。
32-- **moe\_expert\_num** (`int`):必选参数,MoE专家数量,取值范围\[1, 1024\]。32+- **moe\_expert\_num** (`int`):必选参数,MoE专家数量,取值范围\[1, 1024\]。
33- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。33- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
34-- **scales** (`Tensor`):可选参数,表示每个专家的权重,非量化场景不传,动态量化场景可传可不传。若传值要求为3维张量,shape为\(L, shared\_expert\_num + moe\_expert\_num, H\),数据类型支持`float`,数据格式为$ND$,不支持非连续的Tensor。当`quant_mode`为2,`scales`可不为None;当`quant_mode`为0,`scales`必须为None。34+- **scales** (`Tensor`):可选参数,表示每个专家的权重,非量化场景不传,动态量化场景可传可不传。若传值要求为3维张量,shape为\(L, shared\_expert\_num + moe\_expert\_num, H\),数据类型支持`float`,数据格式为$ND$,不支持非连续的Tensor。当`quant_mode`为2,`scales`可不为None;当`quant_mode`为0,`scales`必须为None。
35-- **active\_mask** (`Tensor`):可选参数,表示token是否参与通信。要求是一个2维张量,shape为\(X, BS\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入。默认所有token都会参与通信。35+- **active\_mask** (`Tensor`):可选参数,表示token是否参与通信。要求是一个2维张量,shape为\(X, BS\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入。默认所有token都会参与通信。
36-- **quant\_mode** (`int`):可选参数,表示量化模式。支持取值:0表示非量化(默认),2表示动态量化。36+- **quant\_mode** (`int`):可选参数,表示量化模式。支持取值:0表示非量化(默认),2表示动态量化。
37-- **sync\_flag** (`int`):可选参数,表示同步、异步。支持取值:0表示同步(默认),1表示异步。37+- **sync\_flag** (`int`):可选参数,表示同步、异步。支持取值:0表示同步(默认),1表示异步。
38-- **ffn\_start\_rank\_id** (`int`):可选参数,FFN域起始ID。取值范围\[0, world\_size\),默认为0。38+- **ffn\_start\_rank\_id** (`int`):可选参数,FFN域起始ID。取值范围\[0, world\_size\),默认为0。
39- 
40 39 
41## 返回值说明40## 返回值说明
42 41 
43-- 42+-
44 43 
45## 约束说明44## 约束说明
46 45 
47-- 该接口支持推理场景下使用。46+- 该接口支持推理场景下使用。
48-- 该接口支持静态图模式,分离系列算子必须配套使用。47+- 该接口支持静态图模式,分离系列算子必须配套使用。
49-- 调用接口过程中使用的`group`、`world_size`、`moe_expert_num`参数取值所有卡需保持一致,且网络中不同层中也需保持一致。48+- 调用接口过程中使用的`group`、`world_size`、`moe_expert_num`参数取值所有卡需保持一致,且网络中不同层中也需保持一致。
50-- Atlas A3 训练系列产品/Atlas A3 推理系列产品:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。49+- Atlas A3 训练系列产品/Atlas A3 推理系列产品:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。
51-- 参数里Shape使用的变量如下:50+- 参数里Shape使用的变量如下:
52- - X:表示micro batch sequence size,即token组数,当前版本仅支持 X = 1。51+ - X:表示micro batch sequence size,即token组数,当前版本仅支持 X = 1。
53 52 
54- - H:表示hidden size隐藏层大小,取值为\[1024, 8192\]。53+ - H:表示hidden size隐藏层大小,取值为\[1024, 8192\]。
55 54 
56- - BS:表示batch sequence size,即本卡最终输出的token数量,取值范围为0 < BS ≤ 512。55+ - BS:表示batch sequence size,即本卡最终输出的token数量,取值范围为0 < BS ≤ 512。
57 56 
58- - K:表示选取topK个专家,取值范围为0 < K ≤ 16,同时满足0 < K ≤ moe\_expert\_num。57+ - K:表示选取topK个专家,取值范围为0 < K ≤ 16,同时满足0 < K ≤ moe\_expert\_num。
59 58 
60- - L:表示模型层数,当前版本仅支持 L = 1。59+ - L:表示模型层数,当前版本仅支持 L = 1。
61 60 
62- - shared_expert_num:表示共享专家数量(一个共享专家可以复制部署到多个FFN节点上),取值范围为\[0, 4\]。61+ - shared_expert_num:表示共享专家数量(一个共享专家可以复制部署到多个FFN节点上),取值范围为\[0, 4\]。
63 62 
64-- HCCL通信域缓存区大小:调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。63+- HCCL通信域缓存区大小:调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。
65 64 
66## 调用示例65## 调用示例
67 66 
68-- 单算子模式调用67+- 单算子模式调用
69 68 
70 ```python69 ```python
71 import os70 import os
@@ -221,7 +220,7 @@ torch_npu.npu_attention_to_ffn(x, session_id, micro_batch_id, layer_id, expert_i
221 print("run npu success.")220 print("run npu success.")
222 ```221 ```
223 222 
224-- 图模式调用223+- 图模式调用
225 224 
226 ```python225 ```python
227 # 仅支持静态图226 # 仅支持静态图
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_attention_update.md+3-2
@@ -7,7 +7,6 @@
7| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |7| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |
8| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |8| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |
9 9 
10- 
11## 功能说明10## 功能说明
12 11 
13- API功能:将各SP域PA算子的输出的中间结果`lse``local_out`两个局部变量结果更新成全局结果。12- API功能:将各SP域PA算子的输出的中间结果`lse``local_out`两个局部变量结果更新成全局结果。
@@ -31,7 +30,8 @@
31 $$30 $$
32 31 
33## 函数原型32## 函数原型
34-```33+ 
34+```python
35torch_npu.npu_attention_update(lse, local_out, update_type) -> (Tensor, Tensor)35torch_npu.npu_attention_update(lse, local_out, update_type) -> (Tensor, Tensor)
36```36```
37 37 
@@ -47,6 +47,7 @@ torch_npu.npu_attention_update(lse, local_out, update_type) -> (Tensor, Tensor)
47- **lse_out**(`Tensor`):可选输出,对应公式中的$lse_m$。shape为$(batch \times seqLen \times headNum)$,数据类型为`float32`,数据格式为$ND$。仅当`update_type=1`时有效。47- **lse_out**(`Tensor`):可选输出,对应公式中的$lse_m$。shape为$(batch \times seqLen \times headNum)$,数据类型为`float32`,数据格式为$ND$。仅当`update_type=1`时有效。
48 48 
49## 约束说明49## 约束说明
50+ 
50- 确定性计算:该接口默认为确定性实现,即对于相同的输入,多次执行会产生相同的结果,确保计算结果的可重复性。51- 确定性计算:该接口默认为确定性实现,即对于相同的输入,多次执行会产生相同的结果,确保计算结果的可重复性。
51- 序列并行的并行度SP取值范围[1, 16]。52- 序列并行的并行度SP取值范围[1, 16]。
52- head_dim取值范围[8, 512]且是8的倍数。53- head_dim取值范围[8, 512]且是8的倍数。
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_clipped_swiglu.md+15-14
@@ -9,13 +9,13 @@
9 9 
10## 功能说明10## 功能说明
11 11 
12-- API功能:带截断的Swish门控线性单元激活函数,实现`x`的变体SwiGLU计算。本接口相较于torch_npu.npu_swiglu,新增了部分输入参数:`group_index`、`alpha`、`limit`、`bias`、`interleaved`,用于支持GPT-OSS模型使用的变体SwiGLU以及MoE模型使用的分组场景。12+- API功能:带截断的Swish门控线性单元激活函数,实现`x`的变体SwiGLU计算。本接口相较于torch_npu.npu_swiglu,新增了部分输入参数:`group_index`、`alpha`、`limit`、`bias`、`interleaved`,用于支持GPT-OSS模型使用的变体SwiGLU以及MoE模型使用的分组场景。
13 13 
14-- 计算公式:14+- 计算公式:
15 15 
16 对给定的输入张量`x`,其维度为[a,b,c,d,e,f,g…],进行以下计算:16 对给定的输入张量`x`,其维度为[a,b,c,d,e,f,g…],进行以下计算:
17 17 
18- 1. 将`x`基于输入参数`dim`进行合轴,合轴后维度为[pre, cut, after]。其中cut轴为合轴之后需要切分为两个张量的轴,切分方式分为前后切分或者奇偶切分;pre,after可以等于1。例如当`dim`为3,合轴后`x`的维度为[a * b * c, d, e * f * g * …]。此外,由于after轴的元素为连续存放,且计算操作为逐元素的,因此将cut轴与after轴合并,得到`x`的维度为[pre, cut * after]。18+ 1. 将`x`基于输入参数`dim`进行合轴,合轴后维度为[pre, cut, after]。其中cut轴为合轴之后需要切分为两个张量的轴,切分方式分为前后切分或者奇偶切分;pre,after可以等于1。例如当`dim`为3,合轴后`x`的维度为[a \* b \* c, d, e \* f \* g \*…]。此外,由于after轴的元素为连续存放,且计算操作为逐元素的,因此将cut轴与after轴合并,得到`x`的维度为[pre, cut* after]。
19 19 
20 2. 根据输入参数`group_index`, 对`x`的pre轴进行过滤处理,公式如下:20 2. 根据输入参数`group_index`, 对`x`的pre轴进行过滤处理,公式如下:
21 $$21 $$
@@ -72,22 +72,23 @@
72 72
73## 函数原型73## 函数原型
74 74 
75-```75+```python
76torch_npu.npu_clipped_swiglu(x, *, group_index=None, dim=-1, alpha=1.702, limit=7.0, bias=1.0, interleaved=True) -> Tensor76torch_npu.npu_clipped_swiglu(x, *, group_index=None, dim=-1, alpha=1.702, limit=7.0, bias=1.0, interleaved=True) -> Tensor
77```77```
78 78 
79## 参数说明79## 参数说明
80 80 
81-- **x** (`Tensor`):必选参数,表示目标张量。数据类型支持`float16`、`bfloat16`、`float32`,不支持非连续的`Tensor`,数据格式为$ND$,`x`的维数必须大于1维且第`dim`轴为偶数。81+- **x** (`Tensor`):必选参数,表示目标张量。数据类型支持`float16`、`bfloat16`、`float32`,不支持非连续的`Tensor`,数据格式为$ND$,`x`的维数必须大于1维且第`dim`轴为偶数。
82- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。82- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
83-- **group_index** (`Tensor`):可选参数,表示对`x`进行分组的情况。要求为1维张量,第i个元素代表第i组需要处理的`x`合轴后的token数量,数据类型支持`int64`,数据格式$ND$。默认值为None,表示不对`x`进行分组处理。83+- **group_index** (`Tensor`):可选参数,表示对`x`进行分组的情况。要求为1维张量,第i个元素代表第i组需要处理的`x`合轴后的token数量,数据类型支持`int64`,数据格式$ND$。默认值为None,表示不对`x`进行分组处理。
84-- **dim** (`int`):可选参数,表示需要对`x`进行切分的维度序号,取值范围为[-x.dim(), x.dim()-1],默认值为-1。84+- **dim** (`int`):可选参数,表示需要对`x`进行切分的维度序号,取值范围为[-x.dim(), x.dim()-1],默认值为-1。
85-- **alpha** (`float`):可选参数,表示glu激活函数系数,默认值为1.702。85+- **alpha** (`float`):可选参数,表示glu激活函数系数,默认值为1.702。
86-- **limit** (`float`):可选参数,表示变体SwiGLU输入门限,默认值为7.0。86+- **limit** (`float`):可选参数,表示变体SwiGLU输入门限,默认值为7.0。
87-- **bias** (`float`):可选参数,表示变体SwiGLU计算中的偏差,默认值为1.0。87+- **bias** (`float`):可选参数,表示变体SwiGLU计算中的偏差,默认值为1.0。
88-- **interleaved** (`bool`):可选参数,表示输入`x`是否按奇偶方式切分,True表示为奇偶方式切分,False表示为前后方式切分,默认值为为True。88+- **interleaved** (`bool`):可选参数,表示输入`x`是否按奇偶方式切分,True表示为奇偶方式切分,False表示为前后方式切分,默认值为为True。
89 89 
90## 返回值说明90## 返回值说明
91+ 
91`Tensor`92`Tensor`
92 93 
93代表公式中的`y`,表示激活函数的输出,数据类型同输入`x`,在维度上,第`dim`维是输入`x``1/2`,其余维度与输入`x`相同,数据格式为$ND$。94代表公式中的`y`,表示激活函数的输出,数据类型同输入`x`,在维度上,第`dim`维是输入`x``1/2`,其余维度与输入`x`相同,数据格式为$ND$。
@@ -99,7 +100,7 @@ torch_npu.npu_clipped_swiglu(x, *, group_index=None, dim=-1, alpha=1.702, limit=
99 100 
100## 调用示例101## 调用示例
101 102 
102-- 单算子模式调用103+- 单算子模式调用
103 104 
104 ```python105 ```python
105 import torch106 import torch
@@ -120,7 +121,7 @@ torch_npu.npu_clipped_swiglu(x, *, group_index=None, dim=-1, alpha=1.702, limit=
120 )121 )
121 ```122 ```
122 123 
123-- 图模式调用124+- 图模式调用
124 125 
125 ```python126 ```python
126 import torch127 import torch
@@ -158,4 +159,4 @@ torch_npu.npu_clipped_swiglu(x, *, group_index=None, dim=-1, alpha=1.702, limit=
158 clipped_swiglu_model = ClippedSwigluModel().npu()159 clipped_swiglu_model = ClippedSwigluModel().npu()
159 clipped_swiglu_model = torch.compile(clipped_swiglu_model, backend=npu_backend, dynamic=True)160 clipped_swiglu_model = torch.compile(clipped_swiglu_model, backend=npu_backend, dynamic=True)
160 y = clipped_swiglu_model(x, group_index, -1, 1.702, 7.0, 1.0, True)161 y = clipped_swiglu_model(x, group_index, -1, 1.702, 7.0, 1.0, True)
161- ```162+ ```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_convert_weight_to_int4pack.md+2-3
@@ -13,7 +13,7 @@
13 13 
14## 函数原型14## 函数原型
15 15 
16-```16+```python
17torch_npu.npu_convert_weight_to_int4pack(weight,inner_k_tiles=0) -> Tensor17torch_npu.npu_convert_weight_to_int4pack(weight,inner_k_tiles=0) -> Tensor
18```18```
19 19 
@@ -23,6 +23,7 @@ torch_npu.npu_convert_weight_to_int4pack(weight,inner_k_tiles=0) -> Tensor
23- **inner_k_tiles** (`int`):用于指定内部打包格式中,多少个K-tiles被打包在一起,默认值为`0`**预留参数,暂未使用**23- **inner_k_tiles** (`int`):用于指定内部打包格式中,多少个K-tiles被打包在一起,默认值为`0`**预留参数,暂未使用**
24 24 
25## 返回值说明25## 返回值说明
26+ 
26`Tensor`27`Tensor`
27 28 
28代表`int4`打包后的输出,数据类型为`int32`,shape为$(k, n/8)$, $(n, k/8)$,数据格式支持$ND$。29代表`int4`打包后的输出,数据类型为`int32`,shape为$(k, n/8)$, $(n, k/8)$,数据格式支持$ND$。
@@ -32,7 +33,6 @@ torch_npu.npu_convert_weight_to_int4pack(weight,inner_k_tiles=0) -> Tensor
32- 该接口支持推理场景下使用。33- 该接口支持推理场景下使用。
33- 该接口支持图模式。34- 该接口支持图模式。
34 35 
35- 
36## 调用示例36## 调用示例
37 37 
38- 单算子模式调用38- 单算子模式调用
@@ -144,4 +144,3 @@ torch_npu.npu_convert_weight_to_int4pack(weight,inner_k_tiles=0) -> Tensor
144 [ 34.9375, -15.1797, -23.1094, ..., -13.6797, 8.7734, 6.8750]],144 [ 34.9375, -15.1797, -23.1094, ..., -13.6797, 8.7734, 6.8750]],
145 device='npu:0', dtype=torch.float16)145 device='npu:0', dtype=torch.float16)
146 ```146 ```
147- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_cross_entropy_loss.md+6-8
@@ -9,8 +9,8 @@
9 9 
10## 功能说明10## 功能说明
11 11 
12-- API功能:计算输入`input`和标签`target`之间的交叉熵损失。此API将原生`CrossEntropyLoss`中的log_softmax和nll_loss融合,降低计算时使用的内存。12+- API功能:计算输入`input`和标签`target`之间的交叉熵损失。此API将原生`CrossEntropyLoss`中的log_softmax和nll_loss融合,降低计算时使用的内存。
13-- 计算公式:13+- 计算公式:
14 14 
15 公式中*x*是输入`input`,*y* 是标签`target`,*weight*是权重,*C* 是标签数,*N* 是批处理大小。15 公式中*x*是输入`input`,*y* 是标签`target`,*weight*是权重,*C* 是标签数,*N* 是批处理大小。
16 16 
@@ -33,10 +33,9 @@
33 logProb_{n,c} = x_{n,c} - lse_n33 logProb_{n,c} = x_{n,c} - lse_n
34 $$34 $$
35 35 
36- 
37## 函数原型36## 函数原型
38 37 
39-```38+```python
40torch_npu.npu_cross_entropy_loss(input, target, weight=None, reduction="mean", ignore_index=-100, label_smoothing=0.0, lse_square_scale_for_zloss=0.0, return_zloss=False) -> (Tensor, Tensor, Tensor, Tensor)39torch_npu.npu_cross_entropy_loss(input, target, weight=None, reduction="mean", ignore_index=-100, label_smoothing=0.0, lse_square_scale_for_zloss=0.0, return_zloss=False) -> (Tensor, Tensor, Tensor, Tensor)
41```40```
42 41 
@@ -67,7 +66,8 @@ torch_npu.npu_cross_entropy_loss(input, target, weight=None, reduction="mean", i
67- 输出中仅`loss`支持梯度计算。66- 输出中仅`loss`支持梯度计算。
68 67 
69## 调用示例68## 调用示例
70-- 当reduction设置为`mean`时,示例如下:69+ 
70+- 当reduction设置为`mean`时,示例如下:
71 71 
72 ```python72 ```python
73 import torch73 import torch
@@ -86,8 +86,7 @@ torch_npu.npu_cross_entropy_loss(input, target, weight=None, reduction="mean", i
86 loss.backward()86 loss.backward()
87 ```87 ```
88 88
89- 89+- 当reduction设置为`sum`时,示例如下:
90-- 当reduction设置为`sum`时,示例如下:
91 90 
92 ```python91 ```python
93 import torch92 import torch
@@ -105,4 +104,3 @@ torch_npu.npu_cross_entropy_loss(input, target, weight=None, reduction="mean", i
105 104 
106 loss.backward()105 loss.backward()
107 ```106 ```
108- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_dense_lightning_indexer_grad_kl_loss.md+12-12
@@ -1,6 +1,7 @@
1# torch_npu.npu_dense_lightning_indexer_grad_kl_loss1# torch_npu.npu_dense_lightning_indexer_grad_kl_loss
2 2 
3## 产品支持情况3## 产品支持情况
4+ 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
5|------------------------------| :------: |6|------------------------------| :------: |
6| <term>Atlas A3 训练系列产品</term> | √ |7| <term>Atlas A3 训练系列产品</term> | √ |
@@ -54,7 +55,6 @@
54 dW\mathop{{}}\nolimits_{{t,:}}=dI\mathop{{}}\nolimits_{{t,:}}\text{@} \left( ReLU \left( S\mathop{{}}\nolimits_{{t,:}} \left) \left) \mathop{{}}\nolimits^{\top}\right. \right. \right. \right.55 dW\mathop{{}}\nolimits_{{t,:}}=dI\mathop{{}}\nolimits_{{t,:}}\text{@} \left( ReLU \left( S\mathop{{}}\nolimits_{{t,:}} \left) \left) \mathop{{}}\nolimits^{\top}\right. \right. \right. \right.
55 $$56 $$
56 57
57- 
58 $$58 $$
59 d\mathop{{\tilde{q}}}\nolimits_{{t,:}}=dS\mathop{{}}\nolimits_{{t,:}}@\tilde{K}\mathop{{}}\nolimits_{{:t,:}}59 d\mathop{{\tilde{q}}}\nolimits_{{t,:}}=dS\mathop{{}}\nolimits_{{t,:}}@\tilde{K}\mathop{{}}\nolimits_{{:t,:}}
60 $$60 $$
@@ -72,6 +72,7 @@ npu_dense_lightning_indexer_grad_kl_loss(query, key, query_index, key_index, wei
72```72```
73 73 
74## 参数说明74## 参数说明
75+ 
75**query**(`Tensor`):必选参数,表示Attention中的query,对应公式中的$Q$。数据格式支持$ND$,数据类型支持`bfloat16``float16`。shape支持$(B, S1, N1, D)$、$(T1, N1, D)$。76**query**(`Tensor`):必选参数,表示Attention中的query,对应公式中的$Q$。数据格式支持$ND$,数据类型支持`bfloat16``float16`。shape支持$(B, S1, N1, D)$、$(T1, N1, D)$。
76 77 
77**key**(`Tensor`):必选参数,表示Attention中的key,对应公式中的$K$。数据格式支持$ND$,数据类型支持`bfloat16``float16`。shape支持$(B, S2, N2, D)$、$(T2, N2, D)$。78**key**(`Tensor`):必选参数,表示Attention中的key,对应公式中的$K$。数据格式支持$ND$,数据类型支持`bfloat16``float16`。shape支持$(B, S2, N2, D)$、$(T2, N2, D)$。
@@ -108,19 +109,18 @@ npu_dense_lightning_indexer_grad_kl_loss(query, key, query_index, key_index, wei
108 109 
109**next_tokens**(`int`):可选参数,用于稀疏计算,表示Attention需要和后几个token计算关联。数据类型支持`int64`,默认值2^63-1。110**next_tokens**(`int`):可选参数,用于稀疏计算,表示Attention需要和后几个token计算关联。数据类型支持`int64`,默认值2^63-1。
110 111 
111- 
112- 
113- 
114## 返回值说明112## 返回值说明
115-- **d\_query\_index**(`Tensor`):对应公式中的$d\tilde{Q}$,表示`query_index`的梯度,数据类型支持`bfloat16`、`float16`。113+ 
116-- **d\_key\_index**(`Tensor`):对应公式中的$d\tilde{K}$,表示`key_index`的梯度,数据类型支持`bfloat16`、`float16`。114+- **d\_query\_index**(`Tensor`):对应公式中的$d\tilde{Q}$,表示`query_index`的梯度,数据类型支持`bfloat16`、`float16`。
117-- **d\_weights**(`Tensor`):对应公式中的$dW$,表示`weights`的梯度,数据类型支持`bfloat16`、`float16`、`float32`115+- **d\_key\_index**(`Tensor`):对应公式中的$d\tilde{K}$,表示`key_index`的梯度,数据类型支持`bfloat16`、`float16`。
118-- **loss**(`Tensor`):对应公式中的$Loss$,表示网络正向输出和golden值差异,数据类型支持`float32`。116+- **d\_weights**(`Tensor`):对应公式中的$dW$,表示`weights`梯度,数据类型支持`bfloat16`、`float16`、`float32`。
117+- **loss**(`Tensor`):对应公式中的$Loss$,表示网络正向输出和golden值的差异,数据类型支持`float32`
119 118 
120## 约束说明119## 约束说明
121-- 参数query、key、query_index、key_index的数据类型应保持一致。120+ 
122-- 参数weights不为`float32`时,参数query、key、query_index、key_index、weights的数据类型应保持一致。121+- 参数query、key、query_index、key_index的数据类型应保持一致。
123-- shape值约束:122+-weights不为`float32`时,参数query、key、query_index、key_index、weights的数据类型应保持一致。
123+- shape数值约束:
124 124 
125| 规格项 | 规格 | 规格说明 |125| 规格项 | 规格 | 规格说明 |
126|-----------|------------|-----------------|126|-----------|------------|-----------------|
@@ -133,8 +133,8 @@ npu_dense_lightning_indexer_grad_kl_loss(query, key, query_index, key_index, wei
133| D | 128 | - |133| D | 128 | - |
134| Dr | 64 | - |134| Dr | 64 | - |
135 135 
136- 
137## 调用示例136## 调用示例
137+ 
138- 单算子模式调用138- 单算子模式调用
139 139 
140 ```python140 ```python
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_dense_lightning_indexer_softmax_lse.md+11-10
@@ -1,6 +1,7 @@
1# torch_npu.npu_dense_lightning_indexer_softmax_lse1# torch_npu.npu_dense_lightning_indexer_softmax_lse
2 2 
3## 产品支持情况3## 产品支持情况
4+ 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
5|------------------------------| :------: |6|------------------------------| :------: |
6| <term>Atlas A3 训练系列产品</term> | √ |7| <term>Atlas A3 训练系列产品</term> | √ |
@@ -28,11 +29,12 @@ $maxIndex$,$sumIndex$作为输出传递给接口npu_dense_lightning_indexer_gr
28 29 
29## 函数原型30## 函数原型
30 31 
31-```32+```python
32npu_dense_lightning_indexer_softmax_lse(query_index, key_index, weights, *, actual_seq_qlen=None, actual_seq_klen=None, layout='BSND', sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1) -> (Tensor, Tensor)33npu_dense_lightning_indexer_softmax_lse(query_index, key_index, weights, *, actual_seq_qlen=None, actual_seq_klen=None, layout='BSND', sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1) -> (Tensor, Tensor)
33```34```
34 35 
35## 参数说明36## 参数说明
37+ 
36**query_index**(`Tensor`):必选参数,表示Lightning Indexer正向的输入query,对应公式中的$\tilde{Q}$。数据格式支持$ND$,数据类型支持`bfloat16`、`float16`。shape支持$(B, S1, N1index, D)$、$(T1, N1index, D)$。38**query_index**(`Tensor`):必选参数,表示Lightning Indexer正向的输入query,对应公式中的$\tilde{Q}$。数据格式支持$ND$,数据类型支持`bfloat16`、`float16`。shape支持$(B, S1, N1index, D)$、$(T1, N1index, D)$。
37 39 
38**key_index**(`Tensor`):必选参数,表示Lightning Indexer正向的输入key,对应公式中的$\tilde{K}$。数据格式支持$ND$,数据类型支持`bfloat16`、`float16`。shape支持$(B, S2, N2index, D)$、$(T2, N2index, D)$。40**key_index**(`Tensor`):必选参数,表示Lightning Indexer正向的输入key,对应公式中的$\tilde{K}$。数据格式支持$ND$,数据类型支持`bfloat16`、`float16`。shape支持$(B, S2, N2index, D)$、$(T2, N2index, D)$。
@@ -51,18 +53,17 @@ npu_dense_lightning_indexer_softmax_lse(query_index, key_index, weights, *, actu
51 53 
52**next_tokens**(`int`):可选参数,用于稀疏计算,表示Attention需要和后几个token计算关联。数据类型支持`int64`,默认值2^63-1。54**next_tokens**(`int`):可选参数,用于稀疏计算,表示Attention需要和后几个token计算关联。数据类型支持`int64`,默认值2^63-1。
53 55 
54- 
55- 
56- 
57## 返回值说明56## 返回值说明
58-- **softmax_max_index**(`Tensor`):表示softmax计算使用的max值,对应公式中的$maxIndex$,数据格式支持$ND$,数据类型支持`float32`。57+ 
59-- **softmax_sum_index**(`Tensor`):表示softmax计算使用的sum值,对应公式中的$sumIndex58+- **softmax_max_index**(`Tensor`):表示softmax计算使用的max值,对应公式中的$maxIndex$,数据格式支持$ND$,数据类型支持`float32`。
59+- **softmax_sum_index**(`Tensor`):表示softmax计算使用的sum值,对应公式中的$sumIndex
60$,数据格式支持$ND$,数据类型支持`float32`60$,数据格式支持$ND$,数据类型支持`float32`
61 61 
62## 约束说明62## 约束说明
63-- 参数query_index、key_index的数据类型应保持一致。63+ 
64-- 参数weights不为`float32`时,参数query_index、key_index、weights的数据类型应保持一致。64+- 参数query_index、key_index的数据类型应保持一致。
65-- shape值约束:65+-weights不为`float32`时,参数query_index、key_index、weights的数据类型应保持一致。
66+- shape数值约束:
66 67 
67| 规格项 | 规格 | 规格说明 |68| 规格项 | 规格 | 规格说明 |
68|-----------|------------|-----------------|69|-----------|------------|-----------------|
@@ -72,8 +73,8 @@ $,数据格式支持$ND$,数据类型支持`float32`。
72| N2index | 1 | - |73| N2index | 1 | - |
73| D | 128 | - |74| D | 128 | - |
74 75 
75- 
76## 调用示例76## 调用示例
77+ 
77- 单算子模式调用78- 单算子模式调用
78 79 
79 ```python80 ```python
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_dequant_swiglu_quant.md+52-50
@@ -9,27 +9,27 @@
9 9 
10## 功能说明10## 功能说明
11 11 
12-- API功能:对张量`x`做dequant反量化+swiglu激活+quant量化操作,同时支持分组。12+- API功能:对张量`x`做dequant反量化+swiglu激活+quant量化操作,同时支持分组。
13-- 计算公式:13+- 计算公式:
14- - group分组(目前仅支持count模式):14+ - group分组(目前仅支持count模式):
15 15 
16 对输入`x`进行分组计算,`group_index`表示每个group分组的Tokens数,每组使用不同的量化scale(如`weight_scale`、`activation_scale`、`quant_scale`)。16 对输入`x`进行分组计算,`group_index`表示每个group分组的Tokens数,每组使用不同的量化scale(如`weight_scale`、`activation_scale`、`quant_scale`)。
17 17 
18 举例说明:假设x.shape=\[128, 2H\],group\_index=\[2, 1, 3\],表示有3个group,对应的scale维度为\[3, 2H\]。每个group数据使用不同的scale分别做dequant反量化+swiglu激活+quant量化操作。18 举例说明:假设x.shape=\[128, 2H\],group\_index=\[2, 1, 3\],表示有3个group,对应的scale维度为\[3, 2H\]。每个group数据使用不同的scale分别做dequant反量化+swiglu激活+quant量化操作。
19 19 
20- - group0=x\[0:2, :\],scale0=scale\[0, :\]20+ - group0=x\[0:2, :\],scale0=scale\[0, :\]
21- - group1=x\[2:3, :\],scale1=scale\[1, :\]21+ - group1=x\[2:3, :\],scale1=scale\[1, :\]
22- - group2=x\[3:6, :\],scale2=scale\[2, :\]22+ - group2=x\[3:6, :\],scale2=scale\[2, :\]
23 23 
24- - dequant反量化:分别进行权重量化的反量化、激活量化的反量化。24+ - dequant反量化:分别进行权重量化的反量化、激活量化的反量化。
25 $$25 $$
26 x=x*weight\_scale\\26 x=x*weight\_scale\\
27 x=x*activation\_scale27 x=x*activation\_scale
28 $$28 $$
29 29 
30- - swiglu激活:对反量化后的`x`做swiglu(通过attr activate\_left控制左右激活)。30+ - swiglu激活:对反量化后的`x`做swiglu(通过attr activate\_left控制左右激活)。
31 31 
32- - 当swiglu_mode=0(标准Swiglu),以左激活为例:32+ - 当swiglu_mode=0(标准Swiglu),以左激活为例:
33 $$33 $$
34 swiglu(x)=swish(x[:,0:H])*x[:,H:2H]34 swiglu(x)=swish(x[:,0:H])*x[:,H:2H]
35 $$35 $$
@@ -38,7 +38,7 @@
38 swish(z)=z*sigmoid(z)38 swish(z)=z*sigmoid(z)
39 $$39 $$
40 40 
41- - 当swiglu_mode=1(变种Swiglu),以左激活为例:41+ - 当swiglu_mode=1(变种Swiglu),以左激活为例:
42 $$42 $$
43 x1=clamp(x[:,0:H],-clamp\_limit,clamp\_limit)\\43 x1=clamp(x[:,0:H],-clamp\_limit,clamp\_limit)\\
44 x2=clamp(x[:,H:2H],-clamp\_limit,clamp\_limit)\\44 x2=clamp(x[:,H:2H],-clamp\_limit,clamp\_limit)\\
@@ -49,20 +49,20 @@
49 swish(z,\alpha,bias)=(z+bias)*sigmoid(\alpha*(z+bias))49 swish(z,\alpha,bias)=(z+bias)*sigmoid(\alpha*(z+bias))
50 $$50 $$
51 51 
52- - quant量化:52+ - quant量化:
53- 1. (可选)先进行smooth量化。53+ 1. (可选)先进行smooth量化。
54 $$54 $$
55 out=out*quant\_scale55 out=out*quant\_scale
56 $$56 $$
57 57 
58- 2. 动态/静态量化操作:对激活后的结果进行量化,以动态量化为例。58+ 2. 动态/静态量化操作:对激活后的结果进行量化,以动态量化为例。
59 $$59 $$
60 out,scale=dynamicquant(out)60 out,scale=dynamicquant(out)
61 $$61 $$
62 62 
63## 函数原型63## 函数原型
64 64 
65-```65+```python
66torch_npu.npu_dequant_swiglu_quant(x, *, weight_scale=None, activation_scale=None, bias=None, quant_scale=None, quant_offset=None, group_index=None, activate_left=False, quant_mode=0, swiglu_mode=0, clamp_limit=7.0, float glu_alpha=1.702, float glu_bias=1.0) -> (Tensor, Tensor)66torch_npu.npu_dequant_swiglu_quant(x, *, weight_scale=None, activation_scale=None, bias=None, quant_scale=None, quant_offset=None, group_index=None, activate_left=False, quant_mode=0, swiglu_mode=0, clamp_limit=7.0, float glu_alpha=1.702, float glu_bias=1.0) -> (Tensor, Tensor)
67```67```
68 68 
@@ -70,51 +70,53 @@ torch_npu.npu_dequant_swiglu_quant(x, *, weight_scale=None, activation_scale=Non
70 70 
71> [!NOTE] 71> [!NOTE]
72> Tensor中shape使用的变量说明:72> Tensor中shape使用的变量说明:
73->- TokensNum:表示传输的Tokens数,取值≥0。73+>
74->- H:表示嵌入向量长度,取值\>0。74+>- TokensNum:表示传输Tokens数,取值0。
75->- groupNum:表示group\_index输入的长度,取值\>0。75+>- H:表示向量的长度,取值\>0。
76+>- groupNum:表示group\_index输入的长度,取值\>0。
76 77 
77-- **x** (`Tensor`):必选参数,表示目标张量。要求为2维张量,shape为\[TokensNum, 2H\],尾轴为偶数。数据类型支持`int32`、`bfloat16`,数据格式为$ND$。78+- **x** (`Tensor`):必选参数,表示目标张量。要求为2维张量,shape为\[TokensNum, 2H\],尾轴为偶数。数据类型支持`int32`、`bfloat16`,数据格式为$ND$。
78- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。79- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
79-- **weight\_scale** (`Tensor`):可选参数,表示权重量化对应的反量化系数。要求为2维张量,shape为\[groupNum, 2H\],数据类型支持`float32`,数据格式为$ND$。当`x`为`int32`时,要求该参数非None,表示需要做反量化。80+- **weight\_scale** (`Tensor`):可选参数,表示权重量化对应的反量化系数。要求为2维张量,shape为\[groupNum, 2H\],数据类型支持`float32`,数据格式为$ND$。当`x`为`int32`时,要求该参数非None,表示需要做反量化。
80-- **activation\_scale** (`Tensor`):可选参数,表示pertoken权重量化对应的反量化系数。shape为\[TokensNum, 1\],最后一维为1, 其余与x保持一致。数据类型支持`float32`,数据格式为$ND$。当`x`为`int32`时,要求该参数非None,表示需要做反量化。81+- **activation\_scale** (`Tensor`):可选参数,表示pertoken权重量化对应的反量化系数。shape为\[TokensNum, 1\],最后一维为1, 其余与x保持一致。数据类型支持`float32`,数据格式为$ND$。当`x`为`int32`时,要求该参数非None,表示需要做反量化。
81-- **bias** (`Tensor`):可选参数,表示`x`的偏置变量。数据类型支持`int32`,数据格式为$ND$。group_index为2维的场景下bias必须为None。82+- **bias** (`Tensor`):可选参数,表示`x`的偏置变量。数据类型支持`int32`,数据格式为$ND$。group_index为2维的场景下bias必须为None。
82-- **quant\_scale** (`Tensor`):可选参数,表示smooth量化系数。要求为2维张量,shape为\[groupNum, H\],数据类型支持`float32`、`float16`和`bfloat16`,数据格式为$ND$。83+- **quant\_scale** (`Tensor`):可选参数,表示smooth量化系数。要求为2维张量,shape为\[groupNum, H\],数据类型支持`float32`、`float16`和`bfloat16`,数据格式为$ND$。
83-- **quant\_offset** (`Tensor`):可选参数,表示量化中的偏移项。数据类型支持`float32`、`float16`和`bfloat16`,数据格式为$ND$。`group_index`场景下(非None),该参数不生效为None。84+- **quant\_offset** (`Tensor`):可选参数,表示量化中的偏移项。数据类型支持`float32`、`float16`和`bfloat16`,数据格式为$ND$。`group_index`场景下(非None),该参数不生效为None。
84-- **group\_index** (`Tensor`):可选参数,当前只支持count模式,表示该模式下指定分组的Tokens数(要求非负整数)。要求为1维张量,数据类型支持`int64`,数据格式$ND$。85+- **group\_index** (`Tensor`):可选参数,当前只支持count模式,表示该模式下指定分组的Tokens数(要求非负整数)。要求为1维张量,数据类型支持`int64`,数据格式$ND$。
85-- **activate\_left** (`bool`):可选参数,Swiglu流程中是否进行左激活,默认False。86+- **activate\_left** (`bool`):可选参数,Swiglu流程中是否进行左激活,默认False。
86- - 取True时,out=swish\(split\[x, -1, 2\]\[0\]\)\*split\[x, -1, 2\]\[1\]87+ - 取True时,out=swish\(split\[x, -1, 2\]\[0\]\)\*split\[x, -1, 2\]\[1\]
87- - 取False时,out=swish\(split\[x, -1, 2\]\[1\]\)\*split\[x, -1, 2\]\[0\]88+ - 取False时,out=swish\(split\[x, -1, 2\]\[1\]\)\*split\[x, -1, 2\]\[0\]
89+ 
90+- **quant\_mode** (`int`):可选参数,表示量化类型,默认值为0。0表示静态量化,1表示动态量化。`group_index`场景下(非None),只支持动态量化即`quant_mode`为1。
91+- **swiglu\_mode**(`int`):可选参数,swiglu 计算模式,0 表示传统 swiglu,1 表示变种 swiglu(支持 clamp、alpha、bias)。
92+- **clamp\_limit**(`float`):可选参数,swiglu 输入门限,默认 7.0。
93+- **glu\_alpha**(`float`):可选参数,glu 激活函数系数,默认 1.702。
94+- **glu\_bias**(`float`):可选参数,swiglu 计算中的偏差,默认 1.0。
88 95 
89-- **quant\_mode** (`int`):可选参数,表示量化类型,默认值为0。0表示静态量化,1表示动态量化。`group_index`场景下(非None),只支持动态量化即`quant_mode`为1。
90-- **swiglu\_mode**(`int`):可选参数,swiglu 计算模式,0 表示传统 swiglu,1 表示变种 swiglu(支持 clamp、alpha、bias)。
91-- **clamp\_limit**(`float`):可选参数,swiglu 输入门限,默认 7.0。
92-- **glu\_alpha**(`float`):可选参数,glu 激活函数系数,默认 1.702。
93-- **glu\_bias**(`float`):可选参数,swiglu 计算中的偏差,默认 1.0。
94## 返回值说明96## 返回值说明
95 97 
96-- **out** (`Tensor`):表示量化后的输出tensor。要求是2D的Tensor,shape=\[TokensNum, H\],数据类型支持`int8`,数据格式为$ND$。98+- **out** (`Tensor`):表示量化后的输出tensor。要求是2D的Tensor,shape=\[TokensNum, H\],数据类型支持`int8`,数据格式为$ND$。
97-- **scale** (`Tensor`):表示量化的scale参数。要求是1D的Tensor,shape=\[TokensNum\],数据类型支持`float32`,数据格式为$ND$。99+- **scale** (`Tensor`):表示量化的scale参数。要求是1D的Tensor,shape=\[TokensNum\],数据类型支持`float32`,数据格式为$ND$。
98 100 
99## 约束说明101## 约束说明
100 102 
101-- 该接口支持推理场景下使用。103+- 该接口支持推理场景下使用。
102-- 该接口支持图模式。104+- 该接口支持图模式。
103-- `group_index`场景下(非None)约束说明:105+- `group_index`场景下(非None)约束说明:
104- - `group_index`只支持count模式,需要网络保证`group_index`输入的求和不超过`x`的TokensNum维度,否则会出现越界访问。106+ - `group_index`只支持count模式,需要网络保证`group_index`输入的求和不超过`x`的TokensNum维度,否则会出现越界访问。
105- - H轴有维度大小限制:H≤10496同时64对齐场景;规格不满足场景会进行校验。107+ - H轴有维度大小限制:H≤10496同时64对齐场景;规格不满足场景会进行校验。
106- - 输出`out`和`scale`超过`group_index`总和的部分未进行清理处理,该部分内存为垃圾数据,可能会存在inf/nan异常值,网络使用的时候需要注意影响。108+ - 输出`out`和`scale`超过`group_index`总和的部分未进行清理处理,该部分内存为垃圾数据,可能会存在inf/nan异常值,网络使用的时候需要注意影响。
107-- 当 x 为 int32 时,必须提供 weight_scale。109+- 当 x 为 int32 时,必须提供 weight_scale。
108-- 当 x 为 float16 或 bfloat16 时,weight_scale、activation_scale、bias 必须为 None。110+- 当 x 为 float16 或 bfloat16 时,weight_scale、activation_scale、bias 必须为 None。
109-- x 的最后一维长度必须为偶数。111+- x 的最后一维长度必须为偶数。
110-- 当激活维度不是 x 的最后一维时,group_index 必须为 None。112+- 当激活维度不是 x 的最后一维时,group_index 必须为 None。
111-- 当 group_index 非 None 时,仅支持动态量化(quant_mode=1),且 bias、quant_offset 必须为 None。113+- 当 group_index 非 None 时,仅支持动态量化(quant_mode=1),且 bias、quant_offset 必须为 None。
112-- y 的类型仅支持 int8。114+- y 的类型仅支持 int8。
113-- clamp_limit、glu_alpha、glu_bias 仅在 swiglu_mode=1 时生效。115+- clamp_limit、glu_alpha、glu_bias 仅在 swiglu_mode=1 时生效。
114 116 
115## 调用示例117## 调用示例
116 118 
117-- 单算子模式调用119+- 单算子模式调用
118 120 
119 ```python121 ```python
120 import os122 import os
@@ -157,7 +159,7 @@ torch_npu.npu_dequant_swiglu_quant(x, *, weight_scale=None, activation_scale=Non
157 159 
158 ```160 ```
159 161 
160-- 图模式调用162+- 图模式调用
161 163 
162 ```python164 ```python
163 import os165 import os
@@ -235,4 +237,4 @@ torch_npu.npu_dequant_swiglu_quant(x, *, weight_scale=None, activation_scale=Non
235 if __name__ == "__main__":237 if __name__ == "__main__":
236 run_tests()238 run_tests()
237 239 
238- ```240+ ```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_dynamic_block_quant.md+2-2
@@ -11,7 +11,7 @@
11 11 
12## 功能说明12## 功能说明
13 13 
14-- API功能:对输入张量,通过给定的`row_block_size`和`col_block_size`将输入划分成多个数据块,以数据块为基本粒度进行量化。在每个块中,先计算出当前块对应的量化参数`scale`,并根据`scale`对输入进行量化。输出最终的量化结果,以及每个块的量化参数`scale`。14+- API功能:对输入张量,通过给定的`row_block_size`和`col_block_size`将输入划分成多个数据块,以数据块为基本粒度进行量化。在每个块中,先计算出当前块对应的量化参数`scale`,并根据`scale`对输入进行量化。输出最终的量化结果,以及每个块的量化参数`scale`。
15 15 
16- 计算公式:16- 计算公式:
17 $$17 $$
@@ -30,7 +30,7 @@
30 30 
31## 函数原型31## 函数原型
32 32 
33-```33+```python
34torch_npu.npu_dynamic_block_quant(x, *, min_scale=0.0, round_mode="rint", dst_type=1, row_block_size=1, col_block_size=128) -> (Tensor, Tensor)34torch_npu.npu_dynamic_block_quant(x, *, min_scale=0.0, round_mode="rint", dst_type=1, row_block_size=1, col_block_size=128) -> (Tensor, Tensor)
35```35```
36 36 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_dynamic_quant.md+1-2
@@ -29,7 +29,7 @@
29 29 
30## 函数原型30## 函数原型
31 31 
32-```32+```python
33torch_npu.npu_dynamic_quant(x, *, smooth_scales=None, group_index=None, dst_type=None) ->(Tensor, Tensor)33torch_npu.npu_dynamic_quant(x, *, smooth_scales=None, group_index=None, dst_type=None) ->(Tensor, Tensor)
34```34```
35 35 
@@ -155,4 +155,3 @@ torch_npu.npu_dynamic_quant(x, *, smooth_scales=None, group_index=None, dst_type
155 tensor([[0.0080, 0.0422, 0.0219, 0.0132],155 tensor([[0.0080, 0.0422, 0.0219, 0.0132],
156 [0.0176, 0.0069, 0.0093, 0.0368]], device='npu:0')156 [0.0176, 0.0069, 0.0093, 0.0368]], device='npu:0')
157 ```157 ```
158- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_dynamic_quant_asymmetric.md+18-19
@@ -9,11 +9,11 @@
9 9 
10## 功能说明10## 功能说明
11 11 
12-- API功能:12+- API功能:
13 13 
14 对输入的张量进行动态非对称量化。支持pertoken、pertensor和MoE(Mixture of Experts,混合专家模型)场景。14 对输入的张量进行动态非对称量化。支持pertoken、pertensor和MoE(Mixture of Experts,混合专家模型)场景。
15 15 
16-- 计算公式:16+- 计算公式:
17 17
18 pertoken场景,rowMax、rowMin代表按行取最大值、按行取最小值,此处的“行”对应`x`最后一个维度的数据,即一个token。DST_MAX、DST_MIN分别对应量化后dtype的最大值和最小值,公式如下:18 pertoken场景,rowMax、rowMin代表按行取最大值、按行取最小值,此处的“行”对应`x`最后一个维度的数据,即一个token。DST_MAX、DST_MIN分别对应量化后dtype的最大值和最小值,公式如下:
19 19 
@@ -23,11 +23,11 @@
23 y = \text{round}(\frac{\mathbf{x}}{\text{scale}} + \text{offset})23 y = \text{round}(\frac{\mathbf{x}}{\text{scale}} + \text{offset})
24 $$24 $$
25 25 
26- - 若使用smooth quant,非MoE(Mixture of Experts,混合专家模型)场景下,会引入smooth_scales输入,其形状与x最后一个维度大小一致,在进行量化前,会先令x乘以smooth_scales,再按上述公式进行量化。MoE(Mixture of Experts,混合专家模型)场景下会同时引入smooth_scales和group_index,此时smooth_scales中包含多组smooth向量,按group_index中的数值作用到x的不同行上。具体地,假如x包含m个token,smooth_scales有n行,smooth_scales[0]会作用到x[0:group_index[0]]上,smooth_scales[i]会作用到x[group_index[i-1]: group_index[i]]上,i=1,2, ...,n-1。26+ - 若使用smooth quant,非MoE(Mixture of Experts,混合专家模型)场景下,会引入smooth_scales输入,其形状与x最后一个维度大小一致,在进行量化前,会先令x乘以smooth_scales,再按上述公式进行量化。MoE(Mixture of Experts,混合专家模型)场景下会同时引入smooth_scales和group_index,此时smooth_scales中包含多组smooth向量,按group_index中的数值作用到x的不同行上。具体地,假如x包含m个token,smooth_scales有n行,smooth_scales[0]会作用到x[0:group_index[0]]上,smooth_scales[i]会作用到x[group_index[i-1]: group_index[i]]上,i=1,2, ...,n-1。
27 27 
28## 函数原型28## 函数原型
29 29 
30-```30+```python
31torch_npu.npu_dynamic_quant_asymmetric(x, *, smooth_scales=None, group_index=None, dst_type=None, quant_mode="pertoken") -> (Tensor, Tensor, Tensor)31torch_npu.npu_dynamic_quant_asymmetric(x, *, smooth_scales=None, group_index=None, dst_type=None, quant_mode="pertoken") -> (Tensor, Tensor, Tensor)
32```32```
33 33 
@@ -36,13 +36,13 @@ torch_npu.npu_dynamic_quant_asymmetric(x, *, smooth_scales=None, group_index=Non
36- **x** (`Tensor`):必选参数,需要进行量化的源数据张量,数据类型支持`float16``bfloat16`,数据格式支持ND,支持非连续的Tensor。输入`x`的维度必须大于1。进行int4量化时,要求x形状的最后一维是8的整数倍。36- **x** (`Tensor`):必选参数,需要进行量化的源数据张量,数据类型支持`float16``bfloat16`,数据格式支持ND,支持非连续的Tensor。输入`x`的维度必须大于1。进行int4量化时,要求x形状的最后一维是8的整数倍。
37- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。37- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
38- **smooth_scales** (`Tensor`):可选参数,对`x`进行scales的张量,数据类型支持`float16`、`bfloat16`,数据格式支持$ND$,支持非连续的Tensor。38- **smooth_scales** (`Tensor`):可选参数,对`x`进行scales的张量,数据类型支持`float16`、`bfloat16`,数据格式支持$ND$,支持非连续的Tensor。
39- - 在非MoE场景shape必须是1维,和`x`的最后一维相等。39+ - 在非MoE场景shape必须是1维,和`x`的最后一维相等。
40- - 在MoE场景shape是2维[E, H]。其中E是专家数,取值范围在[1, 1024]且与group_index的第一维相同;H是x的最后一维。40+ - 在MoE场景shape是2维[E, H]。其中E是专家数,取值范围在[1, 1024]且与group_index的第一维相同;H是x的最后一维。
41- - 单算子模式下`smooth_scales`的dtype必须和`x`保持一致,图模式下可以不一致。41+ - 单算子模式下`smooth_scales`的dtype必须和`x`保持一致,图模式下可以不一致。
42- **group_index** (`Tensor`):可选参数,对`smooth_scales`进行分组下标(代表`x`的行数索引),仅在MoE场景下生效。数据类型支持`int32`,数据格式支持$ND$,支持非连续的Tensor。`group_index`的shape为[E,],E的取值范围在[1, 1024]且与smooth_scales第一维相同。tensor的取值必须递增且范围为[1, S],最后一个值必须等于S(S代表输入`x`的行数,是`x`的shape除最后一维度外的乘积)。42- **group_index** (`Tensor`):可选参数,对`smooth_scales`进行分组下标(代表`x`的行数索引),仅在MoE场景下生效。数据类型支持`int32`,数据格式支持$ND$,支持非连续的Tensor。`group_index`的shape为[E,],E的取值范围在[1, 1024]且与smooth_scales第一维相同。tensor的取值必须递增且范围为[1, S],最后一个值必须等于S(S代表输入`x`的行数,是`x`的shape除最后一维度外的乘积)。
43- **dst_type** (`ScalarType`):可选参数,指定量化输出的类型,传None时当作`int8`处理。43- **dst_type** (`ScalarType`):可选参数,指定量化输出的类型,传None时当作`int8`处理。
44- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`int8`、`quint4x2`。44+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`int8`、`quint4x2`。
45- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`int8`、`quint4x2`。45+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`int8`、`quint4x2`。
46- **quant_mode** (`str`):可选参数,量化模式,支持"pertoken"、"pertensor"。默认值为"pertoken"。若`group_index`不为None,只支持"pertoken"。46- **quant_mode** (`str`):可选参数,量化模式,支持"pertoken"、"pertensor"。默认值为"pertoken"。若`group_index`不为None,只支持"pertoken"。
47 47 
48## 返回值说明48## 返回值说明
@@ -53,14 +53,14 @@ torch_npu.npu_dynamic_quant_asymmetric(x, *, smooth_scales=None, group_index=Non
53 53 
54## 约束说明54## 约束说明
55 55 
56-- 该接口支持推理场景下使用。56+- 该接口支持推理场景下使用。
57-- 该接口支持图模式。57+- 该接口支持图模式。
58-- 使用可选参数`smooth_scales`、`group_index`、`dst_type`时,必须使用关键字传参。58+- 使用可选参数`smooth_scales`、`group_index`、`dst_type`时,必须使用关键字传参。
59 59 
60## 调用示例60## 调用示例
61 61 
62-- 单算子模式调用62+- 单算子模式调用
63- - 只有一个输入`x`,进行`int8`量化63+ - 只有一个输入`x`,进行`int8`量化
64 64 
65 ```python65 ```python
66 import torch66 import torch
@@ -70,7 +70,7 @@ torch_npu.npu_dynamic_quant_asymmetric(x, *, smooth_scales=None, group_index=Non
70 print(y, scale, offset)70 print(y, scale, offset)
71 ```71 ```
72 72 
73- - 只有一个输入`x`,进行`int4`量化73+ - 只有一个输入`x`,进行`int4`量化
74 74 
75 ```python75 ```python
76 import torch76 import torch
@@ -80,7 +80,7 @@ torch_npu.npu_dynamic_quant_asymmetric(x, *, smooth_scales=None, group_index=Non
80 print(y, scale, offset)80 print(y, scale, offset)
81 ```81 ```
82 82 
83- - 使用`smooth_scales`输入,非MoE场景(不使用`group_index`),进行`int8`量化83+ - 使用`smooth_scales`输入,非MoE场景(不使用`group_index`),进行`int8`量化
84 84 
85 ```python85 ```python
86 import torch86 import torch
@@ -91,7 +91,7 @@ torch_npu.npu_dynamic_quant_asymmetric(x, *, smooth_scales=None, group_index=Non
91 print(y, scale, offset)91 print(y, scale, offset)
92 ```92 ```
93 93 
94- - 使用`smooth_scales`输入,MoE场景(使用`group_index`),进行`int8`量化94+ - 使用`smooth_scales`输入,MoE场景(使用`group_index`),进行`int8`量化
95 95 
96 ```python96 ```python
97 import torch97 import torch
@@ -103,7 +103,7 @@ torch_npu.npu_dynamic_quant_asymmetric(x, *, smooth_scales=None, group_index=Non
103 print(y, scale, offset)103 print(y, scale, offset)
104 ```104 ```
105 105 
106-- 图模式调用106+- 图模式调用
107 107 
108 ```python108 ```python
109 import torch109 import torch
@@ -135,4 +135,3 @@ torch_npu.npu_dynamic_quant_asymmetric(x, *, smooth_scales=None, group_index=Non
135 print(scale)135 print(scale)
136 print(offset)136 print(offset)
137 ```137 ```
138- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_fast_gelu.md+2-2
@@ -27,7 +27,7 @@
27 27 
28## 函数原型28## 函数原型
29 29 
30-```30+```python
31torch_npu.npu_fast_gelu(input) -> Tensor31torch_npu.npu_fast_gelu(input) -> Tensor
32```32```
33 33 
@@ -41,6 +41,7 @@ torch_npu.npu_fast_gelu(input) -> Tensor
41- <term>Atlas 推理系列产品</term>:数据类型仅支持`float16``float32`41- <term>Atlas 推理系列产品</term>:数据类型仅支持`float16``float32`
42 42 
43## 返回值说明43## 返回值说明
44+ 
44`Tensor`45`Tensor`
45 46 
46代表`fast_gelu`的计算结果。47代表`fast_gelu`的计算结果。
@@ -101,4 +102,3 @@ torch_npu.npu_fast_gelu(input) -> Tensor
101 shape of y: (4, 2048, 16, 128)102 shape of y: (4, 2048, 16, 128)
102 dtype of y: float32103 dtype of y: float32
103 ```104 ```
104- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_ffn.md+2-2
@@ -23,7 +23,7 @@
23 23 
24## 函数原型24## 函数原型
25 25 
26-```26+```python
27torch_npu.npu_ffn(x, weight1, weight2, activation, *, expert_tokens=None, expert_tokens_index=None, bias1=None, bias2=None, scale=None, offset=None, deq_scale1=None, deq_scale2=None, antiquant_scale1=None, antiquant_scale2=None, antiquant_offset1=None, antiquant_offset2=None, inner_precise=None, output_dtype=None) -> Tensor27torch_npu.npu_ffn(x, weight1, weight2, activation, *, expert_tokens=None, expert_tokens_index=None, bias1=None, bias2=None, scale=None, offset=None, deq_scale1=None, deq_scale2=None, antiquant_scale1=None, antiquant_scale2=None, antiquant_offset1=None, antiquant_offset2=None, inner_precise=None, output_dtype=None) -> Tensor
28```28```
29 29 
@@ -65,6 +65,7 @@ torch_npu.npu_ffn(x, weight1, weight2, activation, *, expert_tokens=None, expert
65- **output_dtype** (`ScalarType`):可选参数,该参数只在量化场景生效,其他场景不生效。表示输出Tensor的数据类型,支持输入`float16`、`bfloat16`。默认值为`None`,代表输出Tensor数据类型为`float16`。65- **output_dtype** (`ScalarType`):可选参数,该参数只在量化场景生效,其他场景不生效。表示输出Tensor的数据类型,支持输入`float16`、`bfloat16`。默认值为`None`,代表输出Tensor数据类型为`float16`。
66 66 
67## 返回值说明67## 返回值说明
68+ 
68`Tensor`69`Tensor`
69 70 
70一个Tensor类型的输出,对应公式中的输出$y$,数据类型支持`float16``bfloat16`,数据格式支持$ND$,输出维度与`x`一致。71一个Tensor类型的输出,对应公式中的输出$y$,数据类型支持`float16``bfloat16`,数据格式支持$ND$,输出维度与`x`一致。
@@ -168,4 +169,3 @@ torch_npu.npu_ffn(x, weight1, weight2, activation, *, expert_tokens=None, expert
168 -58.9688]], device='npu:0', dtype=torch.float16)169 -58.9688]], device='npu:0', dtype=torch.float16)
169 170 
170 ```171 ```
171- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_ffn_to_attention.md+17-20
@@ -12,28 +12,28 @@
12 12 
13## 函数原型<a name="zh-cn_topic_0000002203575833_section45077510411"></a>13## 函数原型<a name="zh-cn_topic_0000002203575833_section45077510411"></a>
14 14 
15-```15+```python
16torch_npu.npu_ffn_to_attention(x, session_ids, mirco_batch_ids, token_ids, expert_offsets, actual_token_num, group, world_size,token_info_table_shape, token_data_shape, *, attn_rank_table=None) -> ()16torch_npu.npu_ffn_to_attention(x, session_ids, mirco_batch_ids, token_ids, expert_offsets, actual_token_num, group, world_size,token_info_table_shape, token_data_shape, *, attn_rank_table=None) -> ()
17```17```
18 18 
19## 参数说明<a name="zh-cn_topic_0000002203575833_section112637109429"></a>19## 参数说明<a name="zh-cn_topic_0000002203575833_section112637109429"></a>
20 20 
21-- **x** (`Tensor`):必选参数,表示计算使用的token数据,需根据`sessionIds`来发送给其他卡。要求为2维张量,shape为\(Y, H\),表示有Y个token,数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。21+- **x** (`Tensor`):必选参数,表示计算使用的token数据,需根据`sessionIds`来发送给其他卡。要求为2维张量,shape为\(Y, H\),表示有Y个token,数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。
22-- **session\_ids** (`Tensor`):必选参数,每个token的Attention Worker节点索引,决定每个token要发给哪些Attention Worker节点。要求为1维张量,shape为\(Y, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, attnRankNum-1]。22+- **session\_ids** (`Tensor`):必选参数,每个token的Attention Worker节点索引,决定每个token要发给哪些Attention Worker节点。要求为1维张量,shape为\(Y, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, attnRankNum-1]。
23-- **mirco\_batch\_ids** (`Tensor`):必选参数,表示每个token的microBatch索引,要求为1维张量,shape为\(Y, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, MircoBatchNum-1]。23+- **mirco\_batch\_ids** (`Tensor`):必选参数,表示每个token的microBatch索引,要求为1维张量,shape为\(Y, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, MircoBatchNum-1]。
24-- **token\_ids** (`Tensor`):必选参数,表示每个token在microBatch中的token索引,要求为1维张量,shape为\(Y, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, BS-1]。24+- **token\_ids** (`Tensor`):必选参数,表示每个token在microBatch中的token索引,要求为1维张量,shape为\(Y, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, BS-1]。
25-- **expert\_offsets** (`Tensor`):必选参数,表示每个token在tokenInfoTableShape中PerTokenExpertNum的索引,要求为1维张量,shape为\(Y, \),数据类型支持`in32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, ExpertNumPerToken-1]。25+- **expert\_offsets** (`Tensor`):必选参数,表示每个token在tokenInfoTableShape中PerTokenExpertNum的索引,要求为1维张量,shape为\(Y, \),数据类型支持`in32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, ExpertNumPerToken-1]。
26-- **actual\_token\_num** (`Tensor`):必选参数,表示本卡发送的token总数,要求为1维张量,shape为\(1, \),数据类型支持`in64`,数据格式为$ND$,支持非连续的Tensor。张量里value取值为[0, Y]。26+- **actual\_token\_num** (`Tensor`):必选参数,表示本卡发送的token总数,要求为1维张量,shape为\(1, \),数据类型支持`in64`,数据格式为$ND$,支持非连续的Tensor。张量里value取值为[0, Y]。
27-- **group** (`str`):必选参数,通信域名称,专家并行的通信域。字符串长度范围为\[1,128\)。27+- **group** (`str`):必选参数,通信域名称,专家并行的通信域。字符串长度范围为\[1,128\)。
28-- **world\_size**(`int64`):必选参数,通信域size。取值支持\[2, 768\]。28+- **world\_size**(`int64`):必选参数,通信域size。取值支持\[2, 768\]。
29-- **token\_info\_table\_shape**(`List(int)`):必选参数,Token信息列表大小。包含microBatch的大小(MircoBatchNum)、BatchSize大小(Bs)、以及每个Token对应的Expert数量(ExpertNumPerToken)。29+- **token\_info\_table\_shape**(`List(int)`):必选参数,Token信息列表大小。包含microBatch的大小(MircoBatchNum)、BatchSize大小(Bs)、以及每个Token对应的Expert数量(ExpertNumPerToken)。
30-- **token\_data\_shape**(`List(int)`):必选参数,Token信息列表大小。包含microBatch的大小(MircoBatchNum)、BatchSize大小(Bs)、每个Token对应的Expert数量(ExpertNumPerToken)、以及token和scale长度(HS)。30+- **token\_data\_shape**(`List(int)`):必选参数,Token信息列表大小。包含microBatch的大小(MircoBatchNum)、BatchSize大小(Bs)、每个Token对应的Expert数量(ExpertNumPerToken)、以及token和scale长度(HS)。
31-- **attn\_rank\_table** (`Tensor`):可选参数,映射每一个Attention Worker对应的卡Id,要求为1维张量,shape为\(Y, \),数据类型支持`in32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, attnRankNum-1]。31+- **attn\_rank\_table** (`Tensor`):可选参数,映射每一个Attention Worker对应的卡Id,要求为1维张量,shape为\(Y, \),数据类型支持`in32`,数据格式为$ND$,支持非连续的Tensor。张量里value取值范围为\[0, attnRankNum-1]。
32 32 
33## 约束说明<a name="zh-cn_topic_0000002203575833_section12345537164214"></a>33## 约束说明<a name="zh-cn_topic_0000002203575833_section12345537164214"></a>
34 34 
35-- 调用接口过程中使用的`group`、`world_size`、`token_info_table_shape`、`token_data_shape`参数及`HCCL_BUFFSIZE`参数取值所有卡需保持一致,网络中不同层中也需保持一致。35+- 调用接口过程中使用的`group`、`world_size`、`token_info_table_shape`、`token_data_shape`参数及`HCCL_BUFFSIZE`参数取值所有卡需保持一致,网络中不同层中也需保持一致。
36-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。36+- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。
37 37 
38- 参数说明里shape格式说明:38- 参数说明里shape格式说明:
39 - Y:表示本卡需要分发的最大token数量。39 - Y:表示本卡需要分发的最大token数量。
@@ -56,15 +56,12 @@ torch_npu.npu_ffn_to_attention(x, session_ids, mirco_batch_ids, token_ids, exper
56 56 
57 - sharedExpertNum:表示共享专家数量(一个共享专家可以复制部署到多个ffnRank卡上),取值范围为0 ≤ `sharedExpertNum` ≤ 4。57 - sharedExpertNum:表示共享专家数量(一个共享专家可以复制部署到多个ffnRank卡上),取值范围为0 ≤ `sharedExpertNum` ≤ 4。
58 58 
59- 59+- 通信域使用约束:
60- 
61-- 通信域使用约束:
62 - FFNtoAttention算子的通信域中不允许有其他算子。60 - FFNtoAttention算子的通信域中不允许有其他算子。
63 61 
64- 
65## 调用示例<a name="zh-cn_topic_0000002203575833_section14459801435"></a>62## 调用示例<a name="zh-cn_topic_0000002203575833_section14459801435"></a>
66 63 
67-- 单算子模式调用64+- 单算子模式调用
68 65 
69 ```python66 ```python
70 import os67 import os
@@ -209,7 +206,7 @@ torch_npu.npu_ffn_to_attention(x, session_ids, mirco_batch_ids, token_ids, exper
209 print("run npu success.")206 print("run npu success.")
210 ```207 ```
211 208 
212-- 图模式调用209+- 图模式调用
213 210 
214 ```python211 ```python
215 # 仅支持静态图212 # 仅支持静态图
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_fused_floyd_attention.md+4-1
@@ -19,9 +19,10 @@ $$
19$$19$$
20\text{attention\_out} = \text{einsum}(\text{weights}, \text{value}_1) + \text{einsum}(\text{weights}, \text{value}_2)20\text{attention\_out} = \text{einsum}(\text{weights}, \text{value}_1) + \text{einsum}(\text{weights}, \text{value}_2)
21$$21$$
22+ 
22## 函数原型23## 函数原型
23 24 
24-```25+```python
25torch_npu.npu_fused_floyd_attention(query_ik, key_ij, value_ij, key_jk, value_jk, *, atten_mask=None, scale_value=1.) -> (Tensor, Tensor, Tensor)26torch_npu.npu_fused_floyd_attention(query_ik, key_ij, value_ij, key_jk, value_jk, *, atten_mask=None, scale_value=1.) -> (Tensor, Tensor, Tensor)
26```27```
27 28 
@@ -36,6 +37,7 @@ torch_npu.npu_fused_floyd_attention(query_ik, key_ij, value_ij, key_jk, value_jk
36- **scale_value** (`float`):可选参数,代表缩放系数,对应公式中的$scale\_value$,数据类型支持`float`。默认值为1。37- **scale_value** (`float`):可选参数,代表缩放系数,对应公式中的$scale\_value$,数据类型支持`float`。默认值为1。
37 38 
38## 返回值说明39## 返回值说明
40+ 
39- **softmax_max_out** (`Tensor`):输出张量,Softmax计算的Max中间结果,用于反向计算。数据类型支持`float`,输出的shape类型为[BHNM8]。数据格式支持$ND$。41- **softmax_max_out** (`Tensor`):输出张量,Softmax计算的Max中间结果,用于反向计算。数据类型支持`float`,输出的shape类型为[BHNM8]。数据格式支持$ND$。
40- **softmax_sum_out** (`Tensor`):输出张量,Softmax计算的Sum中间结果,用于反向计算。数据类型支持`float`,输出的shape类型为[BHNM8]。数据格式支持$ND$。42- **softmax_sum_out** (`Tensor`):输出张量,Softmax计算的Sum中间结果,用于反向计算。数据类型支持`float`,输出的shape类型为[BHNM8]。数据格式支持$ND$。
41- **attention_out** (`Tensor`):输出张量,计算公式的最终输出,对应公式中的$attention\_out$。数据类型支持`bfloat16`、`float16`。数据类型和shape类型与`query_ik`保持一致,数据格式支持$ND$,输入shape支持[BHNMD]。43- **attention_out** (`Tensor`):输出张量,计算公式的最终输出,对应公式中的$attention\_out$。数据类型支持`bfloat16`、`float16`。数据类型和shape类型与`query_ik`保持一致,数据格式支持$ND$,输入shape支持[BHNMD]。
@@ -64,6 +66,7 @@ torch_npu.npu_fused_floyd_attention(query_ik, key_ij, value_ij, key_jk, value_jk
64- 支持PyTorch2.6.0及以上版本。66- 支持PyTorch2.6.0及以上版本。
65 67 
66## 调用示例68## 调用示例
69+ 
67```python70```python
68import torch71import torch
69import torch_npu72import torch_npu
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_fused_infer_attention_score.md+184-185
@@ -9,8 +9,8 @@
9 9 
10## 功能说明<a name="zh-cn_topic_0000001832267082_section14441124184110"></a>10## 功能说明<a name="zh-cn_topic_0000001832267082_section14441124184110"></a>
11 11 
12-- API功能:适配增量&全量推理场景的FlashAttention算子,既可以支持全量计算场景(PromptFlashAttention),也可支持增量计算场景(IncreFlashAttention)。当`query`矩阵的S为1,进入IncreFlashAttention分支,其余场景进入PromptFlashAttention分支。12+- API功能:适配增量&全量推理场景的FlashAttention算子,既可以支持全量计算场景(PromptFlashAttention),也可支持增量计算场景(IncreFlashAttention)。当`query`矩阵的S为1,进入IncreFlashAttention分支,其余场景进入PromptFlashAttention分支。
13-- 计算公式:13+- 计算公式:
14 14 
15 $$15 $$
16 attention\_out = softmax \left(scale * (query * key^\top) + atten\_mask \right) * value16 attention\_out = softmax \left(scale * (query * key^\top) + atten\_mask \right) * value
@@ -18,7 +18,7 @@
18 18 
19## 函数原型<a name="zh-cn_topic_0000001832267082_section45077510411"></a>19## 函数原型<a name="zh-cn_topic_0000001832267082_section45077510411"></a>
20 20 
21-```21+```python
22torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None, atten_mask=None, actual_seq_lengths=None, actual_seq_lengths_kv=None, dequant_scale1=None, quant_scale1=None, dequant_scale2=None, quant_scale2=None, quant_offset2=None, antiquant_scale=None, antiquant_offset=None, block_table=None, query_padding_size=None, kv_padding_size=None, key_antiquant_scale=None, key_antiquant_offset=None, value_antiquant_scale=None, value_antiquant_offset=None, key_shared_prefix=None, value_shared_prefix=None, actual_shared_prefix_len=None, query_rope=None, key_rope=None, key_rope_antiquant_scale=None, num_heads=1, scale=1.0, pre_tokens=2147483647, next_tokens=2147483647, input_layout="BSH", num_key_value_heads=0, sparse_mode=0, inner_precise=0, block_size=0, antiquant_mode=0, softmax_lse_flag=False, key_antiquant_mode=0, value_antiquant_mode=0) -> (Tensor, Tensor)22torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None, atten_mask=None, actual_seq_lengths=None, actual_seq_lengths_kv=None, dequant_scale1=None, quant_scale1=None, dequant_scale2=None, quant_scale2=None, quant_offset2=None, antiquant_scale=None, antiquant_offset=None, block_table=None, query_padding_size=None, kv_padding_size=None, key_antiquant_scale=None, key_antiquant_offset=None, value_antiquant_scale=None, value_antiquant_offset=None, key_shared_prefix=None, value_shared_prefix=None, actual_shared_prefix_len=None, query_rope=None, key_rope=None, key_rope_antiquant_scale=None, num_heads=1, scale=1.0, pre_tokens=2147483647, next_tokens=2147483647, input_layout="BSH", num_key_value_heads=0, sparse_mode=0, inner_precise=0, block_size=0, antiquant_mode=0, softmax_lse_flag=False, key_antiquant_mode=0, value_antiquant_mode=0) -> (Tensor, Tensor)
23```23```
24 24 
@@ -37,8 +37,8 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
37 37 
38- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。38- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
39- **pse_shift** (`Tensor`):可选参数。在attention结构内部的位置编码参数,数据类型支持`float16`、`bfloat16`,数据类型与`query`的数据类型需满足数据类型推导规则。不支持非连续的Tensor,数据格式支持$ND$。如不使用该功能时可传入None。39- **pse_shift** (`Tensor`):可选参数。在attention结构内部的位置编码参数,数据类型支持`float16`、`bfloat16`,数据类型与`query`的数据类型需满足数据类型推导规则。不支持非连续的Tensor,数据格式支持$ND$。如不使用该功能时可传入None。
40- - Q_S大于1,要求在`pse_shift`为`float16`类型时,此时的`query`为`float16`或`int8`类型;而在`pse_shift`为`bfloat16`类型时,要求此时的`query`为`bfloat16`类型。输入shape类型需为(B, Q\_N, Q_S, KV_S)或(1, Q\_N, Q_S, KV_S),其中Q_S为`query`的shape中的S,KV_S为`key`和`value`的shape中的S。对于`pse_shift`的KV_S为非32对齐的场景,建议padding到32字节来提高性能,多余部分的填充值不做要求。40+ - Q_S大于1,要求在`pse_shift`为`float16`类型时,此时的`query`为`float16`或`int8`类型;而在`pse_shift`为`bfloat16`类型时,要求此时的`query`为`bfloat16`类型。输入shape类型需为(B, Q\_N, Q_S, KV_S)或(1, Q\_N, Q_S, KV_S),其中Q_S为`query`的shape中的S,KV_S为`key`和`value`的shape中的S。对于`pse_shift`的KV_S为非32对齐的场景,建议padding到32字节来提高性能,多余部分的填充值不做要求。
41- - Q_S为1,要求在`pse_shift`为`float16`类型时,此时的`query`为`float16`类型;而在`pse_shift`为`bfloat16`类型时,要求此时的`query`为`bfloat16`类型。输入shape类型需为(B, Q\_N, 1, KV_S)或(1, Q\_N, 1, KV_S),KV_S为`key`和`value`的shape中的S。对于`pse_shift`的KV_S为非32对齐的场景,建议padding到32字节来提高性能,多余部分的填充值不做要求。41+ - Q_S为1,要求在`pse_shift`为`float16`类型时,此时的`query`为`float16`类型;而在`pse_shift`为`bfloat16`类型时,要求此时的`query`为`bfloat16`类型。输入shape类型需为(B, Q\_N, 1, KV_S)或(1, Q\_N, 1, KV_S),KV_S为`key`和`value`的shape中的S。对于`pse_shift`的KV_S为非32对齐的场景,建议padding到32字节来提高性能,多余部分的填充值不做要求。
42 42 
43- **atten_mask** (`Tensor`):可选参数。对Q(`query`)、K(`key`)的结果进行mask,用于指示是否计算Token间的相关性,数据类型支持`bool`、`int8`和`uint8`。不支持非连续的Tensor,数据格式支持$ND$。如果不使用该功能可传入None。43- **atten_mask** (`Tensor`):可选参数。对Q(`query`)、K(`key`)的结果进行mask,用于指示是否计算Token间的相关性,数据类型支持`bool`、`int8`和`uint8`。不支持非连续的Tensor,数据格式支持$ND$。如果不使用该功能可传入None。
44 - `sparse_mode`为0、1时44 - `sparse_mode`为0、1时
@@ -87,19 +87,19 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
87 87 
88- **num_key_value_heads** (`int`):可选参数。代表`key`、`value`中head个数,用于支持GQA(Grouped-Query Attention,分组查询注意力)场景,数据类型支持`int64`。默认值为0,表示`key`/`value`和`query`的head个数相等,需要满足`num_key_value_heads`整除`num_heads`,`num_heads`与`num_key_value_heads`的比值不能大于64。在BSND、BNSD、BNSD\_BSND(仅支持Q\_S大于1)场景下,还需要与`key`/`value`的N轴shape值相同,否则执行异常。88- **num_key_value_heads** (`int`):可选参数。代表`key`、`value`中head个数,用于支持GQA(Grouped-Query Attention,分组查询注意力)场景,数据类型支持`int64`。默认值为0,表示`key`/`value`和`query`的head个数相等,需要满足`num_key_value_heads`整除`num_heads`,`num_heads`与`num_key_value_heads`的比值不能大于64。在BSND、BNSD、BNSD\_BSND(仅支持Q\_S大于1)场景下,还需要与`key`/`value`的N轴shape值相同,否则执行异常。
89- **sparse_mode** (`int`):可选参数。表示sparse的模式。数据类型支持`int64`。Q\_S为1且不带rope输入时该参数无效。input\_layout为TND、TND\_NTD、NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。89- **sparse_mode** (`int`):可选参数。表示sparse的模式。数据类型支持`int64`。Q\_S为1且不带rope输入时该参数无效。input\_layout为TND、TND\_NTD、NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
90- - `sparse_mode`为0时,代表defaultMask模式,如果`atten_mask`未传入则不做mask操作,忽略`pre_tokens`和`next_tokens`(内部赋值为INT\_MAX);如果传入,则需要传入完整的`atten_mask`矩阵(S1\*S2),表示`pre_tokens`和`next_tokens`之间的部分需要计算。90+ - `sparse_mode`为0时,代表defaultMask模式,如果`atten_mask`未传入则不做mask操作,忽略`pre_tokens`和`next_tokens`(内部赋值为INT\_MAX);如果传入,则需要传入完整的`atten_mask`矩阵(S1\*S2),表示`pre_tokens`和`next_tokens`之间的部分需要计算。
91- - `sparse_mode`为1时,代表allMask,必须传入完整的attenmask矩阵(S1\*S2)。91+ - `sparse_mode`为1时,代表allMask,必须传入完整的attenmask矩阵(S1\*S2)。
92- - `sparse_mode`为2时,代表leftUpCausal模式的mask,需要传入优化后的`atten_mask`矩阵(2048\*2048)。92+ - `sparse_mode`为2时,代表leftUpCausal模式的mask,需要传入优化后的`atten_mask`矩阵(2048\*2048)。
93- - `sparse_mode`为3时,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景,需要传入优化后的`atten_mask`矩阵(2048\*2048)。93+ - `sparse_mode`为3时,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景,需要传入优化后的`atten_mask`矩阵(2048\*2048)。
94- - `sparse_mode`为4时,代表band模式的mask,需要传入优化后的`atten_mask`矩阵(2048\*2048)。94+ - `sparse_mode`为4时,代表band模式的mask,需要传入优化后的`atten_mask`矩阵(2048\*2048)。
95- - `sparse_mode`为5、6、7、8时,分别代表prefix、global、dilated、block\_local,均暂不支持。默认值为0。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。95+ - `sparse_mode`为5、6、7、8时,分别代表prefix、global、dilated、block\_local,均暂不支持。默认值为0。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
96 96 
97- **inner_precise** (`int`):可选参数。一共4种模式:0、1、2、3。一共两位bit位,第0位(bit0)表示高精度或者高性能选择,第1位(bit1)表示是否做行无效修正。数据类型支持`int64`。Q\_S\>1时,`sparse_mode`为0或1,并传入用户自定义mask的情况下,建议开启行无效;Q\_S为1时该参数仅支持`inner_precise`为0和1。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。97- **inner_precise** (`int`):可选参数。一共4种模式:0、1、2、3。一共两位bit位,第0位(bit0)表示高精度或者高性能选择,第1位(bit1)表示是否做行无效修正。数据类型支持`int64`。Q\_S\>1时,`sparse_mode`为0或1,并传入用户自定义mask的情况下,建议开启行无效;Q\_S为1时该参数仅支持`inner_precise`为0和1。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
98 98 
99- - `inner_precise`为0时,代表开启高精度模式,且不做行无效修正。99+ - `inner_precise`为0时,代表开启高精度模式,且不做行无效修正。
100- - `inner_precise`为1时,代表高性能模式,且不做行无效修正。100+ - `inner_precise`为1时,代表高性能模式,且不做行无效修正。
101- - `inner_precise`为2时,代表开启高精度模式,且做行无效修正。101+ - `inner_precise`为2时,代表开启高精度模式,且做行无效修正。
102- - `inner_precise`为3时,代表高性能模式,且做行无效修正。102+ - `inner_precise`为3时,代表高性能模式,且做行无效修正。
103 103 
104 > [!NOTE] 104 > [!NOTE]
105 > `bfloat16`和`int8`不区分高精度和高性能,行无效修正对`float16`、`bfloat16`和`int8`均生效。当前0、1为保留配置值,当计算过程中“参与计算的mask部分”存在某整行全为1的情况时,精度可能会有损失。此时可以尝试将该参数配置为2或3来使能行无效功能以提升精度,但是该配置会导致性能下降。105 > `bfloat16`和`int8`不区分高精度和高性能,行无效修正对`float16`、`bfloat16`和`int8`均生效。当前0、1为保留配置值,当计算过程中“参与计算的mask部分”存在某整行全为1的情况时,精度可能会有损失。此时可以尝试将该参数配置为2或3来使能行无效功能以提升精度,但是该配置会导致性能下降。
@@ -114,12 +114,12 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
114 114 
115 Q\_S大于等于2时仅支持传入值为0、1,Q\_S等于1时支持取值0、1、2、3、4、5。115 Q\_S大于等于2时仅支持传入值为0、1,Q\_S等于1时支持取值0、1、2、3、4、5。
116 116 
117- - `key_antiquant_mode`为0时,代表perchannel模式(perchannel包含pertensor)。117+ - `key_antiquant_mode`为0时,代表perchannel模式(perchannel包含pertensor)。
118- - `key_antiquant_mode`为1时,代表pertoken模式。118+ - `key_antiquant_mode`为1时,代表pertoken模式。
119- - `key_antiquant_mode`为2时,代表pertensor叠加perhead模式。119+ - `key_antiquant_mode`为2时,代表pertensor叠加perhead模式。
120- - `key_antiquant_mode`为3时,代表pertoken叠加perhead模式。120+ - `key_antiquant_mode`为3时,代表pertoken叠加perhead模式。
121- - `key_antiquant_mode`为4时,代表pertoken叠加使用page attention模式管理scale/offset模式。121+ - `key_antiquant_mode`为4时,代表pertoken叠加使用page attention模式管理scale/offset模式。
122- - `key_antiquant_mode`为5时,代表pertoken叠加per head并使用page attention模式管理scale/offset模式。122+ - `key_antiquant_mode`为5时,代表pertoken叠加per head并使用page attention模式管理scale/offset模式。
123 123 
124- **value_antiquant_mode** (`int`):可选参数。表示`value`的伪量化方式,模式编号与`key_antiquant_mode`一致。默认值为0,取值除了`key_antiquant_mode`为0并且`value_antiquant_mode`为1的场景外,需要与`key_antiquant_mode`一致。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。124- **value_antiquant_mode** (`int`):可选参数。表示`value`的伪量化方式,模式编号与`key_antiquant_mode`一致。默认值为0,取值除了`key_antiquant_mode`为0并且`value_antiquant_mode`为1的场景外,需要与`key_antiquant_mode`一致。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
125 125
@@ -127,48 +127,48 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
127 127 
128## 返回值说明<a name="zh-cn_topic_0000001832267082_section22231435517"></a>128## 返回值说明<a name="zh-cn_topic_0000001832267082_section22231435517"></a>
129 129 
130-- **attention\_out** (`Tensor`):公式中的输出,数据类型支持`float16`、`bfloat16`、`int8`。数据格式支持$ND$。限制:该返回值的D维度与`value`的D保持一致,其余维度需要与入参`query`的shape保持一致。130+- **attention\_out** (`Tensor`):公式中的输出,数据类型支持`float16`、`bfloat16`、`int8`。数据格式支持$ND$。限制:该返回值的D维度与`value`的D保持一致,其余维度需要与入参`query`的shape保持一致。
131-- **softmax_lse** (`Tensor`):ring attention算法对query乘key的结果,先取max得到softmax\_max。`query`乘`key`的结果减去softmax\_max,再取exp,最后取sum,得到softmax\_sum,最后对softmax\_sum取log,再加上softmax\_max得到的结果。数据类型支持`float32`,`softmax_lse_flag`为True时,一般情况下,输出shape为\(B, Q\_N, Q\_S, 1\)的Tensor,当input\_layout为TND/NTD\_TND时,输出shape为\(T,Q\_N,1\)的Tensor;`softmax_lse_flag`为False时,则输出shape为\[1\]的值为0的Tensor。131+- **softmax_lse** (`Tensor`):ring attention算法对query乘key的结果,先取max得到softmax\_max。`query`乘`key`的结果减去softmax\_max,再取exp,最后取sum,得到softmax\_sum,最后对softmax\_sum取log,再加上softmax\_max得到的结果。数据类型支持`float32`,`softmax_lse_flag`为True时,一般情况下,输出shape为\(B, Q\_N, Q\_S, 1\)的Tensor,当input\_layout为TND/NTD\_TND时,输出shape为\(T,Q\_N,1\)的Tensor;`softmax_lse_flag`为False时,则输出shape为\[1\]的值为0的Tensor。
132 132 
133## 约束说明<a name="zh-cn_topic_0000001832267082_section12345537164214"></a>133## 约束说明<a name="zh-cn_topic_0000001832267082_section12345537164214"></a>
134 134 
135-- 该接口支持推理场景下使用。135+- 该接口支持推理场景下使用。
136-- 该接口支持图模式。136+- 该接口支持图模式。
137-- 该接口与PyTorch配合使用时,需要保证CANN相关包与PyTorch相关包的版本匹配。137+- 该接口与PyTorch配合使用时,需要保证CANN相关包与PyTorch相关包的版本匹配。
138-- 入参为空的处理:算子内部需要判断参数`query`是否为空,如果是空则直接返回空。参数`query`不为空Tensor,参数`key`、`value`为空Tensor(即S2为0),则`attention_out`按照对应shape大小返回全0。`attention_out`为空Tensor时,返回空。138+- 入参为空的处理:算子内部需要判断参数`query`是否为空,如果是空则直接返回空。参数`query`不为空Tensor,参数`key`、`value`为空Tensor(即S2为0),则`attention_out`按照对应shape大小返回全0。`attention_out`为空Tensor时,返回空。
139-- 参数`key`、`value`中对应tensor的shape需要完全一致;非连续场景下`key`、`value`的tensorlist中的batch只能为1,个数等于`query`的B,N和D需要相等。139+- 参数`key`、`value`中对应tensor的shape需要完全一致;非连续场景下`key`、`value`的tensorlist中的batch只能为1,个数等于`query`的B,N和D需要相等。
140-- `int8`量化相关入参数量与输入、输出数据格式的综合限制:140+- `int8`量化相关入参数量与输入、输出数据格式的综合限制:
141- - 输出为`int8`的场景:入参`dequant_scale1`、`quant_scale1`、`dequant_scale2`、`quant_scale2`需要同时存在,`quant_offset2`可选,不传时默认为0。141+ - 输出为`int8`的场景:入参`dequant_scale1`、`quant_scale1`、`dequant_scale2`、`quant_scale2`需要同时存在,`quant_offset2`可选,不传时默认为0。
142- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:输入为`int8`。142+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:输入为`int8`。
143- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:输入为`int8`。143+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:输入为`int8`。
144 144 
145- - 输出为`float16`的场景:入参`dequant_scale1`、`quant_scale1`、`dequant_scale2`需要同时存在,若存在入参`quant_offset2`或`quant_scale2`(即不为None),则报错并返回。145+ - 输出为`float16`的场景:入参`dequant_scale1`、`quant_scale1`、`dequant_scale2`需要同时存在,若存在入参`quant_offset2`或`quant_scale2`(即不为None),则报错并返回。
146- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:输入为`int8`。146+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:输入为`int8`。
147- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:输入为`int8`。147+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:输入为`int8`。
148 148 
149- - 输入全为`float16`或`bfloat16`,输出为`int8`的场景:入参`quant_scale2`需存在,`quant_offset2`可选,不传时默认为0,若存在入参`dequant_scale1`或`quant_scale1`或`dequant_scale2`(即不为None),则报错并返回。149+ - 输入全为`float16`或`bfloat16`,输出为`int8`的场景:入参`quant_scale2`需存在,`quant_offset2`可选,不传时默认为0,若存在入参`dequant_scale1`或`quant_scale1`或`dequant_scale2`(即不为None),则报错并返回。
150- - 入参`quant_offset2`和`quant_scale2`支持pertensor或perchannel格式,数据类型支持`float32`、`bfloat16`。150+ - 入参`quant_offset2`和`quant_scale2`支持pertensor或perchannel格式,数据类型支持`float32`、`bfloat16`。
151 151 
152-- `antiquant_scale`和`antiquant_offset`参数约束:152+- `antiquant_scale`和`antiquant_offset`参数约束:
153- - 支持perchannel、pertensor和pertoken三种模式:153+ - 支持perchannel、pertensor和pertoken三种模式:
154 154 
155- - perchannel模式:两个参数BNSD场景下shape为\(2, KV\_N, 1, D\),BSND场景下shape为\(2, KV\_N, D\),BSH场景下shape为\(2, H\),N为`num_key_value_heads`。参数数据类型和`query`数据类型相同,`antiquant_mode`置0,当`key`、`value`数据类型为`int8`时支持。155+ - perchannel模式:两个参数BNSD场景下shape为\(2, KV\_N, 1, D\),BSND场景下shape为\(2, KV\_N, D\),BSH场景下shape为\(2, H\),N为`num_key_value_heads`。参数数据类型和`query`数据类型相同,`antiquant_mode`置0,当`key`、`value`数据类型为`int8`时支持。
156- - pertensor模式:两个参数的shape均为\(2,\),数据类型和`query`数据类型相同,`antiquant_mode`置0,当`key`、`value`数据类型为`int8`时支持。156+ - pertensor模式:两个参数的shape均为\(2,\),数据类型和`query`数据类型相同,`antiquant_mode`置0,当`key`、`value`数据类型为`int8`时支持。
157 157 
158- - pertoken模式:两个参数的shape均为\(2, B, KV\_S\),数据类型固定为`float32`,`antiquant_mode`置1,当`key`、`value`数据类型为`int8`时支持。158+ - pertoken模式:两个参数的shape均为\(2, B, KV\_S\),数据类型固定为`float32`,`antiquant_mode`置1,当`key`、`value`数据类型为`int8`时支持。
159 159 
160 算子运行在何种模式根据参数的shape进行判断,dim为1时运行pertensor模式,否则运行perchannel模式。160 算子运行在何种模式根据参数的shape进行判断,dim为1时运行pertensor模式,否则运行perchannel模式。
161 161 
162- - 支持对称量化和非对称量化:162+ - 支持对称量化和非对称量化:
163- - 非对称量化模式下,`antiquant_scale`和`antiquant_offset`参数需同时存在。163+ - 非对称量化模式下,`antiquant_scale`和`antiquant_offset`参数需同时存在。
164- - 对称量化模式下,`antiquant_offset`可以为空(即None);当`antiquant_offset`参数为空时,执行对称量化,否则执行非对称量化。164+ - 对称量化模式下,`antiquant_offset`可以为空(即None);当`antiquant_offset`参数为空时,执行对称量化,否则执行非对称量化。
165 165 
166-- `query_rope`和`key_rope`输入时即为MLA场景,参数约束如下:166+- `query_rope`和`key_rope`输入时即为MLA场景,参数约束如下:
167- - `query_rope`的数据类型、数据格式与`query`一致。167+ - `query_rope`的数据类型、数据格式与`query`一致。
168- - `key_rope`的数据类型、数据格式与`key`一致。168+ - `key_rope`的数据类型、数据格式与`key`一致。
169- - `query_rope`和`key_rope`要求同时配置或同时不配置,不支持只配置其中一个。169+ - `query_rope`和`key_rope`要求同时配置或同时不配置,不支持只配置其中一个。
170- - 当`query_rope`和`key_rope`非空时,`query`的D只支持512、128:170+ - 当`query_rope`和`key_rope`非空时,`query`的D只支持512、128:
171- - 当query的D等于512时:171+ - 当query的D等于512时:
172 - sparse:支持0/3/4;172 - sparse:支持0/3/4;
173 - `query_rope`配置时要求`query`的N为1/2/4/8/16/32/64/128,`query_rope`的shape中D为64,其余维度与`query`一致;173 - `query_rope`配置时要求`query`的N为1/2/4/8/16/32/64/128,`query_rope`的shape中D为64,其余维度与`query`一致;
174 - `key_rope`配置时要求`key`的N为1、D为512,`key_rope`的shape中D为64,其余维度与`key`一致;174 - `key_rope`配置时要求`key`的N为1、D为512,`key_rope`的shape中D为64,其余维度与`key`一致;
@@ -177,30 +177,30 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
177 - 支持开启page attention,此时`block_size`支持16的倍数且不大于1024;177 - 支持开启page attention,此时`block_size`支持16的倍数且不大于1024;
178 - 不支持开启`softmax_lse`、左padding、tensorlist、pse、prefix、伪量化、全量化、后量化。178 - 不支持开启`softmax_lse`、左padding、tensorlist、pse、prefix、伪量化、全量化、后量化。
179 179 
180- - 当query的D等于128时:180+ - 当query的D等于128时:
181 - `input_layout`:BSH、BSND、TND、BNSD、NTD、BSH\_BNSD、BSND\_BNSD、BNSD\_BSND、NTD\_TND。 181 - `input_layout`:BSH、BSND、TND、BNSD、NTD、BSH\_BNSD、BSND\_BNSD、BNSD\_BSND、NTD\_TND。
182 - `query_rope`配置时要求`query_rope`的shape中D为64,其余维度与`query`一致。 182 - `query_rope`配置时要求`query_rope`的shape中D为64,其余维度与`query`一致。
183 - `key_rope`配置时要求`key_rope`的shape中D为64,其余维度与`key`一致。 183 - `key_rope`配置时要求`key_rope`的shape中D为64,其余维度与`key`一致。
184 - 不支持左padding、tensorlist、pse、prefix、伪量化、全量化、后量化。184 - 不支持左padding、tensorlist、pse、prefix、伪量化、全量化、后量化。
185- - 其余约束同TND、NTD\_TND场景下的综合限制保持一致。185+ - 其余约束同TND、NTD\_TND场景下的综合限制保持一致。
186 186 
187- - TND、TND\_NTD、NTD\_TND场景下`query`、`key`、`value`输入的综合限制:187+ - TND、TND\_NTD、NTD\_TND场景下`query`、`key`、`value`输入的综合限制:
188- - `actual_seq_lengths`和`actual_seq_lengths_kv`必须传入,且以该入参元素数量作为Batch值(注意入参元素数量要小于等于4096)。该入参中每个元素的值表示当前Batch与之前所有Batch的Sequence Length和,因此后一个元素的值必须大于等于前一个元素的值;188+ - `actual_seq_lengths`和`actual_seq_lengths_kv`必须传入,且以该入参元素数量作为Batch值(注意入参元素数量要小于等于4096)。该入参中每个元素的值表示当前Batch与之前所有Batch的Sequence Length和,因此后一个元素的值必须大于等于前一个元素的值;
189- - 当query的D等于512时:189+ - 当query的D等于512时:
190- - sparse:支持0/3/4;190+ - sparse:支持0/3/4;
191- - 支持TND、TND\_NTD;191+ - 支持TND、TND\_NTD;
192- - 支持开启page attention,此时`actual_seq_lengths_kv`长度等于`key`/`value`的batch值,代表每个batch的实际长度,值不大于KV\_S;192+ - 支持开启page attention,此时`actual_seq_lengths_kv`长度等于`key`/`value`的batch值,代表每个batch的实际长度,值不大于KV\_S;
193- - 要求`query`的N为1/2/4/8/16/32/64/128,`key`、`value`的N为1;193+ - 要求`query`的N为1/2/4/8/16/32/64/128,`key`、`value`的N为1;
194- - 要求`query_rope`和`key_rope`不等于空,`query_rope`和`key_rope`的D为64;194+ - 要求`query_rope`和`key_rope`不等于空,`query_rope`和`key_rope`的D为64;
195- - 不支持左padding、tensorlist、pse、prefix、伪量化、全量化。195+ - 不支持左padding、tensorlist、pse、prefix、伪量化、全量化。
196 196 
197- - 当query的D不等于512时:197+ - 当query的D不等于512时:
198- - 当`query_rope`和`key_rope`为空时:TND场景,要求Q\_D、K\_D、V\_D等于128,或者Q\_D、K\_D等于192,V\_D等于128/192;NTD场景,不支持V\_D等于192;NTD\_TND场景,要求Q\_D、K\_D等于128/192,V\_D等于128。当`query_rope`和`key_rope`不为空时,要求Q\_D、K\_D、V\_D等于128;198+ - 当`query_rope`和`key_rope`为空时:TND场景,要求Q\_D、K\_D、V\_D等于128,或者Q\_D、K\_D等于192,V\_D等于128/192;NTD场景,不支持V\_D等于192;NTD\_TND场景,要求Q\_D、K\_D等于128/192,V\_D等于128。当`query_rope`和`key_rope`不为空时,要求Q\_D、K\_D、V\_D等于128;
199- - 支持TND、NTD、NTD\_TND;199+ - 支持TND、NTD、NTD\_TND;
200- - page attention场景下仅支持blocksize为16对齐且小于等于1024;200+ - page attention场景下仅支持blocksize为16对齐且小于等于1024;
201- - 不支持左padding、tensorlist、pse、prefix、伪量化、全量化、后量化。201+ - 不支持左padding、tensorlist、pse、prefix、伪量化、全量化、后量化。
202 202 
203-- GQA伪量化场景下KV为NZ格式时的参数约束如下:203+- GQA伪量化场景下KV为NZ格式时的参数约束如下:
204 - 支持perchannel和pertoken模式,`query`数据类型固定为`bfloat16``key`&`value`固定为`int8``query`&`key`&`value`的D仅支持128;query Sequence Length仅支持1-16;204 - 支持perchannel和pertoken模式,`query`数据类型固定为`bfloat16``key`&`value`固定为`int8``query`&`key`&`value`的D仅支持128;query Sequence Length仅支持1-16;
205 - `input_layout`仅支持BSH、BSND、BNSD;205 - `input_layout`仅支持BSH、BSND、BNSD;
206 - 仅支持page_attention场景,blockSize仅支持128或512;206 - 仅支持page_attention场景,blockSize仅支持128或512;
@@ -214,96 +214,96 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
214 - 不支持配置`query_rope``key_rope`214 - 不支持配置`query_rope``key_rope`
215 - 不支持左padding、tensorlist、pse、prefix、后量化;215 - 不支持左padding、tensorlist、pse、prefix、后量化;
216 - num_query_heads与`num_key_value_heads`支持组合有(10, 1)、(64, 8)、(80, 8)、(128, 16)。216 - num_query_heads与`num_key_value_heads`支持组合有(10, 1)、(64, 8)、(80, 8)、(128, 16)。
217-- **当Q\_S大于1时:**217+- **当Q\_S大于1时:**
218- - `query`、`key`、`value`输入,功能使用限制如下:218+ - `query`、`key`、`value`输入,功能使用限制如下:
219- - 支持B轴小于等于65536,D轴32byte不对齐时仅支持到128。219+ - 支持B轴小于等于65536,D轴32byte不对齐时仅支持到128。
220- - 支持N轴小于等于256,支持D轴小于等于512;`input_layout`为BSH或者BSND时,建议N\*D小于65535。220+ - 支持N轴小于等于256,支持D轴小于等于512;`input_layout`为BSH或者BSND时,建议N\*D小于65535。
221- - S支持小于等于20971520(20M)。部分长序列场景下,如果计算量过大可能会导致PFA算子执行超时(aicore error类型报错,errorStr为timeout or trap error),此场景下建议做S切分处理(注:这里计算量会受B、S、N、D等的影响,值越大计算量越大),典型的会超时的长序列(即B、S、N、D的乘积较大)场景包括但不限于:221+ - S支持小于等于20971520(20M)。部分长序列场景下,如果计算量过大可能会导致PFA算子执行超时(aicore error类型报错,errorStr为timeout or trap error),此场景下建议做S切分处理(注:这里计算量会受B、S、N、D等的影响,值越大计算量越大),典型的会超时的长序列(即B、S、N、D的乘积较大)场景包括但不限于:
222- - B=1,Q\_N=20,Q\_S=2097152,D=256,KV\_N=1,KV\_S=2097152。222+ - B=1,Q\_N=20,Q\_S=2097152,D=256,KV\_N=1,KV\_S=2097152。
223- - B=1,Q\_N=2,Q\_S=20971520,D=256,KV\_N=2,KV\_S=20971520。223+ - B=1,Q\_N=2,Q\_S=20971520,D=256,KV\_N=2,KV\_S=20971520。
224- - B=20,Q\_N=1,Q\_S=2097152,D=256,KV\_N=1,KV\_S=2097152。224+ - B=20,Q\_N=1,Q\_S=2097152,D=256,KV\_N=1,KV\_S=2097152。
225- - B=1,Q\_N=10,Q\_S=2097152,D=512,KV\_N=1,KV\_S=2097152。225+ - B=1,Q\_N=10,Q\_S=2097152,D=512,KV\_N=1,KV\_S=2097152。
226 226 
227- - `query`、`key`、`value`输入类型包含`int8`时,D轴需要32对齐;输入类型全为`float16`、`bfloat16`时,D轴需16对齐。227+ - `query`、`key`、`value`输入类型包含`int8`时,D轴需要32对齐;输入类型全为`float16`、`bfloat16`时,D轴需16对齐。
228- - D轴限制:228+ - D轴限制:
229- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`query`、`key`、`value`输入类型包含`int8`时,D轴需要32对齐;`query`、`key`、`value`或`attention_out`类型包含`int4`时,D轴需要64对齐;输入类型全为`float16`、`bfloat16`时,D轴需16对齐。229+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`query`、`key`、`value`输入类型包含`int8`时,D轴需要32对齐;`query`、`key`、`value`或`attention_out`类型包含`int4`时,D轴需要64对齐;输入类型全为`float16`、`bfloat16`时,D轴需16对齐。
230 230 
231- - `actual_seq_lengths`:231+ - `actual_seq_lengths`:
232 232
233 <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于`query`中对应batch的Sequence Length。seqlen的传入长度为1时,每个Batch使用相同seqlen;传入长度大于等于Batch时取seqlen的前Batch个数。其他长度不支持。当`query`的`input_layout`为TND/NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。233 <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于`query`中对应batch的Sequence Length。seqlen的传入长度为1时,每个Batch使用相同seqlen;传入长度大于等于Batch时取seqlen的前Batch个数。其他长度不支持。当`query`的`input_layout`为TND/NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
234 234
235- - `actual_seq_lengths_kv`:235+ - `actual_seq_lengths_kv`:
236 236
237 <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于`key`/`value`中对应batch的Sequence Length。seqlenKv的传入长度为1时,每个Batch使用相同seqlenKv;传入长度大于等于Batch时取seqlenKv的前Batch个数。其他长度不支持。当`key`/`value`的`input_layout`为TND/NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。237 <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于`key`/`value`中对应batch的Sequence Length。seqlenKv的传入长度为1时,每个Batch使用相同seqlenKv;传入长度大于等于Batch时取seqlenKv的前Batch个数。其他长度不支持。当`key`/`value`的`input_layout`为TND/NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
238 238
239- - 参数`sparse_mode`当前仅支持值为0、1、2、3、4的场景,取其它值时会报错。239+ - 参数`sparse_mode`当前仅支持值为0、1、2、3、4的场景,取其它值时会报错。
240 240
241- - `sparse_mode`为0时,`atten_mask`如果为None,或者在左padding场景传入`atten_mask`,则忽略入参pre\_tokens、next\_tokens(内部赋值为INT\_MAX)。241+ - `sparse_mode`为0时,`atten_mask`如果为None,或者在左padding场景传入`atten_mask`,则忽略入参pre\_tokens、next\_tokens(内部赋值为INT\_MAX)。
242- - `sparse_mode`为2、3、4时,`atten_mask`的shape需要为\(S, S\)或\(1, S, S\)或\(1, 1, S, S\),其中S的值需要固定为2048,且需要用户保证传入的`atten_mask`为下三角,不传入`atten_mask`或者传入的shape不正确报错。 242+ - `sparse_mode`为2、3、4时,`atten_mask`的shape需要为\(S, S\)或\(1, S, S\)或\(1, 1, S, S\),其中S的值需要固定为2048,且需要用户保证传入的`atten_mask`为下三角,不传入`atten_mask`或者传入的shape不正确报错。
243- - `sparse_mode`为1、2、3的场景忽略入参pre\_tokens、next\_tokens并按照相关规则赋值。243+ - `sparse_mode`为1、2、3的场景忽略入参pre\_tokens、next\_tokens并按照相关规则赋值。
244 244
245- - kvCache反量化的合成参数场景仅支持`int8`反量化到`float16`。入参`key`、`value`的data range与入参`antiquant_scale`的data range乘积范围在(-1, 1)内,高性能模式可以保证精度,否则需要开启高精度模式来保证精度。245+ - kvCache反量化的合成参数场景仅支持`int8`反量化到`float16`。入参`key`、`value`的data range与入参`antiquant_scale`的data range乘积范围在(-1, 1)内,高性能模式可以保证精度,否则需要开启高精度模式来保证精度。
246- - page attention场景:246+ - page attention场景:
247- - page attention的使能必要条件是`block_table`存在且有效,同时`key`、`value`是按照`block_table`中的索引在一片连续内存中排布,支持`key`、`value`数据类型为`float16`、`bfloat16`。在该场景下`key`、`value`的`input_layout`参数无效。`block_table`中填充的是blockid,当前不会对blockid的合法性进行校验,需用户自行保证。247+ - page attention的使能必要条件是`block_table`存在且有效,同时`key`、`value`是按照`block_table`中的索引在一片连续内存中排布,支持`key`、`value`数据类型为`float16`、`bfloat16`。在该场景下`key`、`value`的`input_layout`参数无效。`block_table`中填充的是blockid,当前不会对blockid的合法性进行校验,需用户自行保证。
248- - `block_size`是用户自定义的参数,该参数的取值会影响page attention的性能,在使能page attention场景下,`block_size`最小为128,最大为512,且要求是128的倍数。通常情况下,page attention可以提高吞吐量,但会带来性能上的下降。248+ - `block_size`是用户自定义的参数,该参数的取值会影响page attention的性能,在使能page attention场景下,`block_size`最小为128,最大为512,且要求是128的倍数。通常情况下,page attention可以提高吞吐量,但会带来性能上的下降。
249 249
250- - page attention场景下,当输入kv cache排布格式为(blocknum, blocksize, H),且KV\_N\*D超过65535时,受硬件指令约束,会被拦截报错。可通过使能GQA(减小KV\_N)或调整kv cache排布格式为(blocknum, KV\_N, blocksize, D)解决。当`query`的`input_layout`为BNSD、TND时,kv cache排布支持(blocknum, blocksize, H)和(blocknum, KV\_N, blocksize, D)两种格式,当`query`的`input_layout`为BSH、BSND时,kv cache排布只支持(blocknum, blocksize, H)一种格式。blocknum不能小于根据`actual_seq_lengths_kv`和`block_size`计算的每个batch的block数量之和。且`key`和`value`的shape需保证一致。250+ - page attention场景下,当输入kv cache排布格式为(blocknum, blocksize, H),且KV\_N\*D超过65535时,受硬件指令约束,会被拦截报错。可通过使能GQA(减小KV\_N)或调整kv cache排布格式为(blocknum, KV\_N, blocksize, D)解决。当`query`的`input_layout`为BNSD、TND时,kv cache排布支持(blocknum, blocksize, H)和(blocknum, KV\_N, blocksize, D)两种格式,当`query`的`input_layout`为BSH、BSND时,kv cache排布只支持(blocknum, blocksize, H)一种格式。blocknum不能小于根据`actual_seq_lengths_kv`和`block_size`计算的每个batch的block数量之和。且`key`和`value`的shape需保证一致。
251- - page attention不支持伪量化场景,不支持tensorlist场景,不支持左padding场景。251+ - page attention不支持伪量化场景,不支持tensorlist场景,不支持左padding场景。
252- - page attention场景下,必须传入`actual_seq_lengths_kv`。252+ - page attention场景下,必须传入`actual_seq_lengths_kv`。
253- - page attention场景下,`block_table`必须为二维,第一维长度需等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为不同batch中最大`actual_seq_lengths_kv`对应的block数量)。253+ - page attention场景下,`block_table`必须为二维,第一维长度需等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为不同batch中最大`actual_seq_lengths_kv`对应的block数量)。
254- - page attention场景下,支持两种格式和`float32`/`bfloat16`,不支持输入`query`为`int8`的场景。254+ - page attention场景下,支持两种格式和`float32`/`bfloat16`,不支持输入`query`为`int8`的场景。
255- - page attention使能场景下,以下场景输入需满足KV\_S\>=maxBlockNumPerSeq\*blockSize:255+ - page attention使能场景下,以下场景输入需满足KV\_S\>=maxBlockNumPerSeq\*blockSize:
256- - 传入`atten_mask`时,如mask shape为(B, 1, Q\_S, KV\_S)。256+ - 传入`atten_mask`时,如mask shape为(B, 1, Q\_S, KV\_S)。
257- - 传入`pse_shift`时,如pseShift shape为(B, Q\_N, Q\_S, KV\_S)。257+ - 传入`pse_shift`时,如pseShift shape为(B, Q\_N, Q\_S, KV\_S)。
258 258
259- - `query`左padding场景:259+ - `query`左padding场景:
260- - `query`左padding场景`query`的搬运起点计算公式为:Q\_S-query\_padding\_size-actual\_seq\_lengths。`query`的搬运终点计算公式为:Q\_S-query\_padding\_size。其中`query`的搬运起点不能小于0,终点不能大于Q\_S,否则结果将不符合预期。260+ - `query`左padding场景`query`的搬运起点计算公式为:Q\_S-query\_padding\_size-actual\_seq\_lengths。`query`的搬运终点计算公式为:Q\_S-query\_padding\_size。其中`query`的搬运起点不能小于0,终点不能大于Q\_S,否则结果将不符合预期。
261- - `query`左padding场景`kv_padding_size`小于0时将被置为0。261+ - `query`左padding场景`kv_padding_size`小于0时将被置为0。
262- - `query`左padding场景需要与`actual_seq_lengths`参数一起使能,否则默认为`query`右padding场景。 262+ - `query`左padding场景需要与`actual_seq_lengths`参数一起使能,否则默认为`query`右padding场景。
263- - `query`左padding场景不支持page attention,不能与`block_table`参数一起使能。263+ - `query`左padding场景不支持page attention,不能与`block_table`参数一起使能。
264 264
265- - kv左padding场景:265+ - kv左padding场景:
266- - kv左padding场景`key`和`value`的搬运起点计算公式为:KV\_S-kv\_padding\_size-actual\_seq\_lengths\_kv。`key`和`value`的搬运终点计算公式为:KV\_S-kv\_padding\_size。其中`key`和`value`的搬运起点不能小于0,终点不能大于KV\_S,否则结果将不符合预期。266+ - kv左padding场景`key`和`value`的搬运起点计算公式为:KV\_S-kv\_padding\_size-actual\_seq\_lengths\_kv。`key`和`value`的搬运终点计算公式为:KV\_S-kv\_padding\_size。其中`key`和`value`的搬运起点不能小于0,终点不能大于KV\_S,否则结果将不符合预期。
267- - kv左padding场景`kv_padding_size`小于0时将被置为0。267+ - kv左padding场景`kv_padding_size`小于0时将被置为0。
268- - kv左padding场景需要与`actual_seq_lengths_kv`参数一起使能,否则默认为kv右padding场景。 268+ - kv左padding场景需要与`actual_seq_lengths_kv`参数一起使能,否则默认为kv右padding场景。
269- - kv左padding场景不支持page attention,不能与`block_table`参数一起使能。269+ - kv左padding场景不支持page attention,不能与`block_table`参数一起使能。
270 270
271- - 入参`quant_scale2`和`quant_offset2`支持pertensor、perchannel量化,支持`float32`、`bfloat16`类型。若传入`quant_offset2`,需保证其类型和shape信息与 `quant_scale2`一致。当输入为`bfloat16`时,同时支持`float32`和`bfloat16`,否则仅支持`float32`。perchannel场景下,当输出layout为BSH时,要求`quant_scale2`所有维度的乘积等于H;其他layout要求乘积等于N\*D。当输出layout为BSH时,`quant_scale2` shape建议传入\(1, 1, H\)或\(H,\);当输出layout为BNSD时,建议传入\(1, Q\_N, 1, D\)或\(Q\_N, D\);当输出为BSND时,建议传入\(1, 1, Q\_N, D\)或\(Q\_N, D)。271+ - 入参`quant_scale2`和`quant_offset2`支持pertensor、perchannel量化,支持`float32`、`bfloat16`类型。若传入`quant_offset2`,需保证其类型和shape信息与 `quant_scale2`一致。当输入为`bfloat16`时,同时支持`float32`和`bfloat16`,否则仅支持`float32`。perchannel场景下,当输出layout为BSH时,要求`quant_scale2`所有维度的乘积等于H;其他layout要求乘积等于N\*D。当输出layout为BSH时,`quant_scale2` shape建议传入\(1, 1, H\)或\(H,\);当输出layout为BNSD时,建议传入\(1, Q\_N, 1, D\)或\(Q\_N, D\);当输出为BSND时,建议传入\(1, 1, Q\_N, D\)或\(Q\_N, D)。
272- - 输出为`int8`,`quant_scale2`和`quant_offset2`为perchannel时,暂不支持左padding、Ring Attention或者D非32Byte对齐的场景。272+ - 输出为`int8`,`quant_scale2`和`quant_offset2`为perchannel时,暂不支持左padding、Ring Attention或者D非32Byte对齐的场景。
273- - 输出为`int8`时,暂不支持sparse为band且preTokens/nextTokens为负数。273+ - 输出为`int8`时,暂不支持sparse为band且preTokens/nextTokens为负数。
274- - `pse_shift`功能使用限制如下:274+ - `pse_shift`功能使用限制如下:
275 275
276- - 支持`query`数据类型为`float16`或`bfloat16`或`int8`场景下使用该功能。276+ - 支持`query`数据类型为`float16`或`bfloat16`或`int8`场景下使用该功能。
277- - `query`、`key`、`value`数据类型为`float16`且`pse_shift`存在时,强制走高精度模式,对应的限制继承自高精度模式的限制。277+ - `query`、`key`、`value`数据类型为`float16`且`pse_shift`存在时,强制走高精度模式,对应的限制继承自高精度模式的限制。
278- - Q\_S需大于等于`query`的S长度,KV\_S需大于等于`key`的S长度。prefix场景KV\_S需大于等于`actual_shared_prefix_len`与`key`的S长度之和。278+ - Q\_S需大于等于`query`的S长度,KV\_S需大于等于`key`的S长度。prefix场景KV\_S需大于等于`actual_shared_prefix_len`与`key`的S长度之和。
279 279
280- - 输出为`int8`,入参`quant_offset2`传入非None和非空tensor值,并且`sparse_mode`、`pre_tokens`和`next_tokens`满足以下条件,矩阵会存在某几行不参与计算的情况,导致计算结果误差,该场景会拦截:280+ - 输出为`int8`,入参`quant_offset2`传入非None和非空tensor值,并且`sparse_mode`、`pre_tokens`和`next_tokens`满足以下条件,矩阵会存在某几行不参与计算的情况,导致计算结果误差,该场景会拦截:
281- - `sparse_mode`为0,`atten_mask`如果非None,每个batch actual\_seq\_lengths-actual\_seq\_lengths\_kv-pre\_tokens\>0或next\_tokens<0时,满足拦截条件。281+ - `sparse_mode`为0,`atten_mask`如果非None,每个batch actual\_seq\_lengths-actual\_seq\_lengths\_kv-pre\_tokens\>0或next\_tokens<0时,满足拦截条件。
282- - `sparse_mode`为1或2,不会出现满足拦截条件的情况。282+ - `sparse_mode`为1或2,不会出现满足拦截条件的情况。
283- - `sparse_mode`为3,每个batch actual\_seq\_lengths\_kv-actual\_seq\_lengths<0,满足拦截条件。283+ - `sparse_mode`为3,每个batch actual\_seq\_lengths\_kv-actual\_seq\_lengths<0,满足拦截条件。
284- - `sparse_mode`为4,pre\_tokens<0或每个batch next\_tokens+actual\_seq\_lengths\_kv-actual\_seq\_lengths<0时,满足拦截条件。284+ - `sparse_mode`为4,pre\_tokens<0或每个batch next\_tokens+actual\_seq\_lengths\_kv-actual\_seq\_lengths<0时,满足拦截条件。
285 285
286- - prefix相关参数约束:286+ - prefix相关参数约束:
287- - `key_shared_prefix`和`value_shared_prefix`要么都为空,要么都不为空。287+ - `key_shared_prefix`和`value_shared_prefix`要么都为空,要么都不为空。
288- - `key_shared_prefix`和`value_shared_prefix`都不为空时,`key_shared_prefix`、`value_shared_prefix`、`key`、`value`的维度相同、dtype保持一致。288+ - `key_shared_prefix`和`value_shared_prefix`都不为空时,`key_shared_prefix`、`value_shared_prefix`、`key`、`value`的维度相同、dtype保持一致。
289- - `key_shared_prefix`和`value_shared_prefix`都不为空时,`key_shared_prefix`的shape第一维batch必须为1,layout为BNSD和BSND情况下N、D轴要与`key`一致、BSH情况下H要与`key`一致,`value_shared_prefix`同理。`key_shared_prefix`和`value_shared_prefix`的S应相等。289+ - `key_shared_prefix`和`value_shared_prefix`都不为空时,`key_shared_prefix`的shape第一维batch必须为1,layout为BNSD和BSND情况下N、D轴要与`key`一致、BSH情况下H要与`key`一致,`value_shared_prefix`同理。`key_shared_prefix`和`value_shared_prefix`的S应相等。
290- - 当`actual_shared_prefix_len`存在时,`actual_shared_prefix_len`的shape需要为\[1\],值不能大于`key_shared_prefix`和`value_shared_prefix`的S。290+ - 当`actual_shared_prefix_len`存在时,`actual_shared_prefix_len`的shape需要为\[1\],值不能大于`key_shared_prefix`和`value_shared_prefix`的S。
291- - 公共前缀的S加上`key`或`value`的S的结果,要满足原先`key`或`value`的S的限制。291+ - 公共前缀的S加上`key`或`value`的S的结果,要满足原先`key`或`value`的S的限制。
292- - prefix不支持page attention场景、不支持左padding场景、不支持tensorlist场景。292+ - prefix不支持page attention场景、不支持左padding场景、不支持tensorlist场景。
293- - prefix场景不支持`query`、`key`、`value`数据类型同时为`int8`。293+ - prefix场景不支持`query`、`key`、`value`数据类型同时为`int8`。
294- - prefix场景,sparse为0或1时,如果传入attenmask,则S2需大于等于`actual_shared_prefix_len`与`key`的S长度之和。294+ - prefix场景,sparse为0或1时,如果传入attenmask,则S2需大于等于`actual_shared_prefix_len`与`key`的S长度之和。
295- - prefix场景,不支持输入qkv全部为`int8`的场景。295+ - prefix场景,不支持输入qkv全部为`int8`的场景。
296 296
297- - kv伪量化参数分离:297+ - kv伪量化参数分离:
298- - 当伪量化参数和KV分离量化参数同时传入时,以KV分离量化参数为准。 298+ - 当伪量化参数和KV分离量化参数同时传入时,以KV分离量化参数为准。
299- - `key_antiquant_mode`和`value_antiquant_mode`取值需要保持一致。299+ - `key_antiquant_mode`和`value_antiquant_mode`取值需要保持一致。
300- - `key_antiquant_scale`和`value_antiquant_scale`要么都为空,要么都不为空;`key\_antiquant_offset`和`value_antiquant_offset`要么都为空,要么都不为空。300+ - `key_antiquant_scale`和`value_antiquant_scale`要么都为空,要么都不为空;`key\_antiquant_offset`和`value_antiquant_offset`要么都为空,要么都不为空。
301- - `key_antiquant_scale`和`value_antiquant_scale`都不为空时,其shape需要保持一致;`key_antiquant_offset`和`value_antiquant_offset`都不为空时,其shape需要保持一致。 301+ - `key_antiquant_scale`和`value_antiquant_scale`都不为空时,其shape需要保持一致;`key_antiquant_offset`和`value_antiquant_offset`都不为空时,其shape需要保持一致。
302- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:302+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
303- - 仅支持pertoken和perchannel模式,pertoken模式下要求两个参数的shape均为\(B, KV\_S\),数据类型固定为`float32`;perchannel模式下要求两个参数的shape为(KV\_N,D),\(KV\_N, D\),\(H\),数据类型固定为`bfloat16`。303+ - 仅支持pertoken和perchannel模式,pertoken模式下要求两个参数的shape均为\(B, KV\_S\),数据类型固定为`float32`;perchannel模式下要求两个参数的shape为(KV\_N,D),\(KV\_N, D\),\(H\),数据类型固定为`bfloat16`。
304- - `key_antiquant_scale`与`value_antiquant_scale`非空场景,要求`query`的s小于等于16;要求`query`的dtype为`bfloat16`,`key`、`value`的dtype为`int8`,输出的dtype为`bfloat16`;不支持tensorlist、左padding、page attention、prefix特性。304+ - `key_antiquant_scale`与`value_antiquant_scale`非空场景,要求`query`的s小于等于16;要求`query`的dtype为`bfloat16`,`key`、`value`的dtype为`int8`,输出的dtype为`bfloat16`;不支持tensorlist、左padding、page attention、prefix特性。
305 305
306- - 管理scale/offset的量化模式如下:306+ - 管理scale/offset的量化模式如下:
307 307
308 > [!NOTE] 308 > [!NOTE]
309 > 注意scale、offset两个参数指`key_antiquant_scale`、`value_antiquant_scale`、`key_antiquant_offset`、`value_antiquant_offset`参数。309 > 注意scale、offset两个参数指`key_antiquant_scale`、`value_antiquant_scale`、`key_antiquant_offset`、`value_antiquant_offset`参数。
@@ -338,53 +338,53 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
338 </tbody>338 </tbody>
339 </table>339 </table>
340 340
341-- **当Q\_S等于1时:**341+- **当Q\_S等于1时:**
342- - `query`、`key`、`value`输入,功能使用限制如下:342+ - `query`、`key`、`value`输入,功能使用限制如下:
343- - 支持B轴小于等于65536,支持N轴小于等于256,支持S轴小于等于262144,支持D轴小于等于512。343+ - 支持B轴小于等于65536,支持N轴小于等于256,支持S轴小于等于262144,支持D轴小于等于512。
344- - `query`、`key`、`value`输入类型均为`int8`的场景暂不支持。344+ - `query`、`key`、`value`输入类型均为`int8`的场景暂不支持。
345- - 在`int4`(`int32`)伪量化场景下,PyTorch入图调用仅支持KV `int4`拼接成`int32`输入(建议通过dynamicQuant生成`int4`格式的数据,因为dynamicQuant就是一个`int32`包括8个`int4`)。345+ - 在`int4`(`int32`)伪量化场景下,PyTorch入图调用仅支持KV `int4`拼接成`int32`输入(建议通过dynamicQuant生成`int4`格式的数据,因为dynamicQuant就是一个`int32`包括8个`int4`)。
346- - 在`int4`(`int32`)伪量化场景下,若KV `int4`拼接成`int32`输入,那么KV的N、D或者H是实际值的八分之一(prefix同理)。并且,`int4`伪量化仅支持D 64对齐(`int32`支持D 8对齐)。346+ - 在`int4`(`int32`)伪量化场景下,若KV `int4`拼接成`int32`输入,那么KV的N、D或者H是实际值的八分之一(prefix同理)。并且,`int4`伪量化仅支持D 64对齐(`int32`支持D 8对齐)。
347 347 
348- - `actual_seq_lengths`:348+ - `actual_seq_lengths`:
349 349
350- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当`query`的`input_layout`不为TND时,Q\_S为1时该参数无效。当`query`的`input_layout`为TND/TND\_NTD时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。350+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当`query`的`input_layout`不为TND时,Q\_S为1时该参数无效。当`query`的`input_layout`为TND/TND\_NTD时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
351 351
352- - `actual_seq_lengths_kv`:352+ - `actual_seq_lengths_kv`:
353 353
354- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于`key`/`value`中对应batch的Sequence Length。seqlenKv的传入长度为1时,每个Batch使用相同seqlenKv;传入长度大于等于Batch时取seqlenKv的前Batch个数。其他长度不支持。当`key`/`value`的`input_layout`为TND/TND\_NTD时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。354+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于`key`/`value`中对应batch的Sequence Length。seqlenKv的传入长度为1时,每个Batch使用相同seqlenKv;传入长度大于等于Batch时取seqlenKv的前Batch个数。其他长度不支持。当`key`/`value`的`input_layout`为TND/TND\_NTD时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
355 355
356- - page attention场景:356+ - page attention场景:
357- - 使能必要条件是`block_table`存在且有效,同时`key`、`value`是按照`block_table`中的索引在一片连续内存中排布,在该场景下`key`、`value`的`input_layout`参数无效。357+ - 使能必要条件是`block_table`存在且有效,同时`key`、`value`是按照`block_table`中的索引在一片连续内存中排布,在该场景下`key`、`value`的`input_layout`参数无效。
358- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:358+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
359- - 支持`key`、`value`数据类型为`float16`、`bfloat16`、`int8`。359+ - 支持`key`、`value`数据类型为`float16`、`bfloat16`、`int8`。
360- - 不支持`query`为`bfloat16`、`float16`,且`key`和`value`为`int4`(`int32`)的场景。360+ - 不支持`query`为`bfloat16`、`float16`,且`key`和`value`为`int4`(`int32`)的场景。
361 361
362- - 该场景下,`block_size`是用户自定义的参数,该参数的取值会影响page attention的性能,`block_size`需要传入非0值,且最大不超过512,`key`、`value`输入类型为`float16`、`bfloat16`时需要16对齐,`key`、`value`输入类型为`int8`时需要32对齐,推荐使用128。通常情况下,page attention可以提高吞吐量,但会带来性能上的下降。362+ - 该场景下,`block_size`是用户自定义的参数,该参数的取值会影响page attention的性能,`block_size`需要传入非0值,且最大不超过512,`key`、`value`输入类型为`float16`、`bfloat16`时需要16对齐,`key`、`value`输入类型为`int8`时需要32对齐,推荐使用128。通常情况下,page attention可以提高吞吐量,但会带来性能上的下降。
363- - 参数`key`、`value`各自对应tensor的shape所有维度相乘不能超过`int32`的表示范围。363+ - 参数`key`、`value`各自对应tensor的shape所有维度相乘不能超过`int32`的表示范围。
364- - page attention场景下,`block_table`必须为二维,第一维长度需等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为不同batch中最大`actual_seq_lengths_kv`对应的block数量)。364+ - page attention场景下,`block_table`必须为二维,第一维长度需等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为不同batch中最大`actual_seq_lengths_kv`对应的block数量)。
365- - page attention场景下,当`query`的`input_layout`为BNSD、TND时,kv cache排布支持(blocknum, blocksize, H)和(blocknum, KV\_N, blocksize, D)两种格式,当`query`的`input_layout`为BSH、BSND时,kv cache排布只支持(blocknum, blocksize, H)一种格式。blocknum不能小于根据`actual_seq_lengths_kv`和`block_size`计算的每个batch的block数量之和。且`key`和`value`的shape需保证一致。365+ - page attention场景下,当`query`的`input_layout`为BNSD、TND时,kv cache排布支持(blocknum, blocksize, H)和(blocknum, KV\_N, blocksize, D)两种格式,当`query`的`input_layout`为BSH、BSND时,kv cache排布只支持(blocknum, blocksize, H)一种格式。blocknum不能小于根据`actual_seq_lengths_kv`和`block_size`计算的每个batch的block数量之和。且`key`和`value`的shape需保证一致。
366- - page attention场景下,kv cache排布为(blocknum, KV\_N, blocksize, D)时性能通常优于kv cache排布为(blocknum, blocksize, H)时的性能,建议优先选择(blocknum, KV\_N, blocksize, D)格式。366+ - page attention场景下,kv cache排布为(blocknum, KV\_N, blocksize, D)时性能通常优于kv cache排布为(blocknum, blocksize, H)时的性能,建议优先选择(blocknum, KV\_N, blocksize, D)格式。
367- - page attention场景下,当输入kv cache排布格式为(blocknum, blocksize, H),且numKvHeads \* headDim 超过64k时,受硬件指令约束,会被拦截报错。可通过使能GQA(减小 numKvHeads)或调整kv cache排布格式为(blocknum, numKvHeads, blocksize, D)解决。367+ - page attention场景下,当输入kv cache排布格式为(blocknum, blocksize, H),且numKvHeads \* headDim 超过64k时,受硬件指令约束,会被拦截报错。可通过使能GQA(减小 numKvHeads)或调整kv cache排布格式为(blocknum, numKvHeads, blocksize, D)解决。
368- - page attention不支持tensorlist场景,不支持左padding场景。368+ - page attention不支持tensorlist场景,不支持左padding场景。
369- - page attention场景的参数`key`、`value`各自对应tensor的shape所有维度相乘不能超过`int32`的表示范围。369+ - page attention场景的参数`key`、`value`各自对应tensor的shape所有维度相乘不能超过`int32`的表示范围。
370- - page attention场景下,使能`atten_mask`,当`sparse_mode`不为2、3、4时,传入的`atten_mask`的最后一维需要大于等于`block_table`的第二维 * `block_size`。370+ - page attention场景下,使能`atten_mask`,当`sparse_mode`不为2、3、4时,传入的`atten_mask`的最后一维需要大于等于`block_table`的第二维 * `block_size`。
371- - page attention场景下,使能`pse_shift`,传入的`pse_shift`的最后一维需要大于等于`block_table`的第二维 * `block_size`。371+ - page attention场景下,使能`pse_shift`,传入的`pse_shift`的最后一维需要大于等于`block_table`的第二维 * `block_size`。
372- - page attention场景下,以下场景输入S需要大于等于`block_table`的第二维 * `block_size`。372+ - page attention场景下,以下场景输入S需要大于等于`block_table`的第二维 * `block_size`。
373- - 使能伪量化pertoken模式:输入参数`antiqunant_scale`和`antiquant_offset`的shape均为\(2, B, S\)。373+ - 使能伪量化pertoken模式:输入参数`antiqunant_scale`和`antiquant_offset`的shape均为\(2, B, S\)。
374- - 使能pertoken叠加perhead模式:两个参数的shape均为\(B, N, S\),数据类型固定为`float32`。支持`key`、`value`数据类型为`int8`、`int4`\(`int32`\)。374+ - 使能pertoken叠加perhead模式:两个参数的shape均为\(B, N, S\),数据类型固定为`float32`。支持`key`、`value`数据类型为`int8`、`int4`\(`int32`\)。
375 375
376- - kv左padding场景:376+ - kv左padding场景:
377- - kvCache的搬运起点计算公式为:Smax-kv\_padding\_size-actual\_seq\_lengths。kvCache的搬运终点计算公式为:Smax-kv\_padding\_size。其中kvCache的搬运起点或终点小于0时,返回数据结果为全0。377+ - kvCache的搬运起点计算公式为:Smax-kv\_padding\_size-actual\_seq\_lengths。kvCache的搬运终点计算公式为:Smax-kv\_padding\_size。其中kvCache的搬运起点或终点小于0时,返回数据结果为全0。
378- - `kv_padding_size`小于0时将被置为0。378+ - `kv_padding_size`小于0时将被置为0。
379- - 使能需要同时存在`actual_seq_lengths`参数,否则默认为kv右padding场景。379+ - 使能需要同时存在`actual_seq_lengths`参数,否则默认为kv右padding场景。
380- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:kv左padding场景不支持Q为`bfloat16`/`float16`、KV为`int4`(`int32`)的场景。380+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:kv左padding场景不支持Q为`bfloat16`/`float16`、KV为`int4`(`int32`)的场景。
381 381
382- - kv伪量化参数分离:382+ - kv伪量化参数分离:
383- - 除了`keyantiquant_mode`为0并且`value_antiquant_mode`为1的场景外,`key_antiquant_mode`和`value_antiquant_mode`取值需要保持一致。 383+ - 除了`keyantiquant_mode`为0并且`value_antiquant_mode`为1的场景外,`key_antiquant_mode`和`value_antiquant_mode`取值需要保持一致。
384- - `key_antiquant_scale`和`value_antiquant_scale`要么都为空,要么都不为空;`key_antiquant_offset`和`value_antiquant_offset`要么都为空,要么都不为空。384+ - `key_antiquant_scale`和`value_antiquant_scale`要么都为空,要么都不为空;`key_antiquant_offset`和`value_antiquant_offset`要么都为空,要么都不为空。
385- - `key_antiquant_scale`和`value_antiquant_scale`都不为空时,除了`key_antiquant_mode`为0并且`value_antiquant_mode`为1的场景外,其shape需要保持一致;`key_antiquant_offset`和`value_antiquant_offset`都不为空时,除了`key_antiquant_mode`为0并且`value_antiquant_mode`为1的场景外,其shape需要保持一致。385+ - `key_antiquant_scale`和`value_antiquant_scale`都不为空时,除了`key_antiquant_mode`为0并且`value_antiquant_mode`为1的场景外,其shape需要保持一致;`key_antiquant_offset`和`value_antiquant_offset`都不为空时,除了`key_antiquant_mode`为0并且`value_antiquant_mode`为1的场景外,其shape需要保持一致。
386- - `int4`(`int32`)伪量化场景不支持后量化。386+ - `int4`(`int32`)伪量化场景不支持后量化。
387- - 管理scale/offset的量化模式如下:387+ - 管理scale/offset的量化模式如下:
388 388
389 > [!NOTE] 389 > [!NOTE]
390 > 注意scale、offset两个参数指`key_antiquant_scale`、`value_antiquant_scale`、`key_antiquant_offset`、`value_antiquant_offset`参数。390 > 注意scale、offset两个参数指`key_antiquant_scale`、`value_antiquant_scale`、`key_antiquant_offset`、`value_antiquant_offset`参数。
@@ -483,16 +483,16 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
483 <tr id="zh-cn_topic_0000001832267082_row194748261012"><td class="cellrowborder" valign="top" headers="mcps1.1.5.1.1 "><p id="zh-cn_topic_0000001832267082_p1154111491113"><a name="zh-cn_topic_0000001832267082_p1154111491113"></a><a name="zh-cn_topic_0000001832267082_p1154111491113"></a>对于value支持pertoken,两个参数的shape均为(1, B, KV_S)并且数据类型固定为float32。</p>483 <tr id="zh-cn_topic_0000001832267082_row194748261012"><td class="cellrowborder" valign="top" headers="mcps1.1.5.1.1 "><p id="zh-cn_topic_0000001832267082_p1154111491113"><a name="zh-cn_topic_0000001832267082_p1154111491113"></a><a name="zh-cn_topic_0000001832267082_p1154111491113"></a>对于value支持pertoken,两个参数的shape均为(1, B, KV_S)并且数据类型固定为float32。</p>
484 </td>484 </td>
485 </tr>485 </tr>
486- </tbody>486+ </tbody>
487 </table>487 </table>
488 488
489- - `pse_shift`功能使用限制如下:489+ - `pse_shift`功能使用限制如下:
490- - `pse_shift`数据类型需与`query`数据类型保持一致。490+ - `pse_shift`数据类型需与`query`数据类型保持一致。
491- - 仅支持D轴对齐,即D轴可以被16整除。491+ - 仅支持D轴对齐,即D轴可以被16整除。
492 492 
493## 调用示例<a name="zh-cn_topic_0000001832267082_section14459801435"></a>493## 调用示例<a name="zh-cn_topic_0000001832267082_section14459801435"></a>
494 494 
495-- 单算子模式调用495+- 单算子模式调用
496 496 
497 ```python497 ```python
498 import torch498 import torch
@@ -522,7 +522,7 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
522 device='npu:0', dtype=torch.float16)522 device='npu:0', dtype=torch.float16)
523 ```523 ```
524 524 
525-- 图模式调用525+- 图模式调用
526 526 
527 ```python527 ```python
528 # 入图方式528 # 入图方式
@@ -586,4 +586,3 @@ torch_npu.npu_fused_infer_attention_score(query, key, value, *, pse_shift=None,
586 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],586 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],
587 device='npu:0', dtype=torch.float16) torch.Size([1, 8, 164, 128])587 device='npu:0', dtype=torch.float16) torch.Size([1, 8, 164, 128])
588 ```588 ```
589- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_fused_infer_attention_score_v2.md+187-187
@@ -1,6 +1,7 @@
1# torch\_npu.npu\_fused\_infer\_attention\_score\_v2<a name="ZH-CN_TOPIC_0000001979260729"></a>1# torch\_npu.npu\_fused\_infer\_attention\_score\_v2<a name="ZH-CN_TOPIC_0000001979260729"></a>
2 2 
3## 产品支持情况 <a name="zh-cn_topic_0000001832267082_section14441124184110"></a>3## 产品支持情况 <a name="zh-cn_topic_0000001832267082_section14441124184110"></a>
4+ 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
5| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
6|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
@@ -10,8 +11,8 @@
10 11 
11## 功能说明<a name="zh-cn_topic_0000001832267082_section14441124184110"></a>12## 功能说明<a name="zh-cn_topic_0000001832267082_section14441124184110"></a>
12 13 
13-- API功能:适配增量&全量推理场景的FlashAttention算子,既可以支持全量计算场景(PromptFlashAttention),也可支持增量计算场景(IncreFlashAttention)。当不涉及system prefix、左padding、kv量化参数合一、pertensor全量化的场景,推荐使用本接口,否则使用老接口`npu_fused_infer_attention_score`。14+- API功能:适配增量&全量推理场景的FlashAttention算子,既可以支持全量计算场景(PromptFlashAttention),也可支持增量计算场景(IncreFlashAttention)。当不涉及system prefix、左padding、kv量化参数合一、pertensor全量化的场景,推荐使用本接口,否则使用老接口`npu_fused_infer_attention_score`。
14-- 计算公式:15+- 计算公式:
15 16 
16 $$17 $$
17 Attention(Q,K,V)=Softmax(\frac{QK^T}{\sqrt{d}})V18 Attention(Q,K,V)=Softmax(\frac{QK^T}{\sqrt{d}})V
@@ -21,7 +22,7 @@
21 22 
22## 函数原型<a name="zh-cn_topic_0000001832267082_section45077510411"></a>23## 函数原型<a name="zh-cn_topic_0000001832267082_section45077510411"></a>
23 24 
24-```25+```python
25torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=None, key_rope=None, pse_shift=None, atten_mask=None, actual_seq_qlen=None, actual_seq_kvlen=None, block_table=None, dequant_scale_query=None, dequant_scale_key=None, dequant_offset_key=None, dequant_scale_value=None, dequant_offset_value=None, dequant_scale_key_rope=None, quant_scale_out=None, quant_offset_out=None, learnable_sink=None, num_query_heads=1, num_key_value_heads=0, softmax_scale=1.0, pre_tokens=2147483647, next_tokens=2147483647, input_layout="BSH", sparse_mode=0, block_size=0, query_quant_mode=0, key_quant_mode=0, value_quant_mode=0, inner_precise=0, return_softmax_lse=False, query_dtype=None, key_dtype=None, value_dtype=None, query_rope_dtype=None, key_rope_dtype=None, key_shared_prefix_dtype=None, value_shared_prefix_dtype=None, dequant_scale_query_dtype=None, dequant_scale_key_dtype=None, dequant_scale_value_dtype=None, dequant_scale_key_rope_dtype=None) -> (Tensor, Tensor)26torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=None, key_rope=None, pse_shift=None, atten_mask=None, actual_seq_qlen=None, actual_seq_kvlen=None, block_table=None, dequant_scale_query=None, dequant_scale_key=None, dequant_offset_key=None, dequant_scale_value=None, dequant_offset_value=None, dequant_scale_key_rope=None, quant_scale_out=None, quant_offset_out=None, learnable_sink=None, num_query_heads=1, num_key_value_heads=0, softmax_scale=1.0, pre_tokens=2147483647, next_tokens=2147483647, input_layout="BSH", sparse_mode=0, block_size=0, query_quant_mode=0, key_quant_mode=0, value_quant_mode=0, inner_precise=0, return_softmax_lse=False, query_dtype=None, key_dtype=None, value_dtype=None, query_rope_dtype=None, key_rope_dtype=None, key_shared_prefix_dtype=None, value_shared_prefix_dtype=None, dequant_scale_query_dtype=None, dequant_scale_key_dtype=None, dequant_scale_value_dtype=None, dequant_scale_key_rope_dtype=None) -> (Tensor, Tensor)
26```27```
27 28 
@@ -32,21 +33,21 @@ torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=No
32> - query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示隐藏层的大小、N(Head Num)表示多头数、D(Head Dim)表示隐藏层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。33> - query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示隐藏层的大小、N(Head Num)表示多头数、D(Head Dim)表示隐藏层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。
33> - Q_S和S1表示query shape中的S,KV_S和S2表示key和value shape中的S,Q_N表示num\_query\_heads,KV_N表示num\_key\_value\_heads。34> - Q_S和S1表示query shape中的S,KV_S和S2表示key和value shape中的S,Q_N表示num\_query\_heads,KV_N表示num\_key\_value\_heads。
34 35 
35-- **query**(`Tensor`):必选参数,表示attention结构的Query输入,对应公式中的`Q`。不支持非连续的Tensor,数据类型支持`float16`、`bfloat16`,数据格式支持ND。 36+- **query**(`Tensor`):必选参数,表示attention结构的Query输入,对应公式中的`Q`。不支持非连续的Tensor,数据类型支持`float16`、`bfloat16`,数据格式支持ND。
36 37
37-- **key**(`Tensor`):必选参数,表示attention结构的Key输入,对应公式中的`K`。不支持非连续的Tensor,数据类型支持`float16`、`bfloat16`、`int8`、`int4`(`int32`),数据格式支持ND。 38+- **key**(`Tensor`):必选参数,表示attention结构的Key输入,对应公式中的`K`。不支持非连续的Tensor,数据类型支持`float16`、`bfloat16`、`int8`、`int4`(`int32`),数据格式支持ND。
38 39
39-- **value**(`Tensor`):必选参数,表示attention结构的Value输入,对应公式中的`V`。不支持非连续的Tensor,数据类型支持`float16`、`bfloat16`、`int8`、`int4`(`int32`),数据格式支持ND。 40+- **value**(`Tensor`):必选参数,表示attention结构的Value输入,对应公式中的`V`。不支持非连续的Tensor,数据类型支持`float16`、`bfloat16`、`int8`、`int4`(`int32`),数据格式支持ND。
40 41
41- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。42- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
42-- **query\_rope**(`Tensor`):可选参数,表示MLA(Multi-head Latent Attention)结构中`query`的rope信息,数据类型支持`float16`、`bfloat16`,不支持非连续的Tensor,数据格式支持ND。43+- **query\_rope**(`Tensor`):可选参数,表示MLA(Multi-head Latent Attention)结构中`query`的rope信息,数据类型支持`float16`、`bfloat16`,不支持非连续的Tensor,数据格式支持ND。
43-- **key\_rope**(`Tensor`):可选参数,表示MLA(Multi-head Latent Attention)结构中的`key`的rope信息,数据类型支持`float16`、`bfloat16`,不支持非连续的Tensor,数据格式支持ND。44+- **key\_rope**(`Tensor`):可选参数,表示MLA(Multi-head Latent Attention)结构中的`key`的rope信息,数据类型支持`float16`、`bfloat16`,不支持非连续的Tensor,数据格式支持ND。
44-- **pse\_shift**(`Tensor`):可选参数,表示attention结构内部的位置编码参数,数据类型支持`float16`、`bfloat16`,数据类型与`query`数据类型需满足类型推导规则。不支持非连续的Tensor,数据格式支持ND。如不使用该功能可传入None。45+- **pse\_shift**(`Tensor`):可选参数,表示attention结构内部的位置编码参数,数据类型支持`float16`、`bfloat16`,数据类型与`query`数据类型需满足类型推导规则。不支持非连续的Tensor,数据格式支持ND。如不使用该功能可传入None。
45 46 
46- - Q\_S大于1,当`pse_shift`为`float16`类型时,要求`query`为float16或int8类型;当`pse_shift`为`bfloat16`类型时,要求`query`为`bfloat16`类型。输入shape类型需为\(B, Q\_N, Q\_S, KV\_S\)或\(1, Q\_N, Q\_S, KV\_S\)。对于`pse_shift`的KV\_S为非32对齐的场景,建议padding到32字节来提高性能,多余部分的填充值不做要求。47+ - Q\_S大于1,当`pse_shift`为`float16`类型时,要求`query`为float16或int8类型;当`pse_shift`为`bfloat16`类型时,要求`query`为`bfloat16`类型。输入shape类型需为\(B, Q\_N, Q\_S, KV\_S\)或\(1, Q\_N, Q\_S, KV\_S\)。对于`pse_shift`的KV\_S为非32对齐的场景,建议padding到32字节来提高性能,多余部分的填充值不做要求。
47- - Q\_S为1,当`pse_shift`为`float16`类型时,要求`query`为`float16`类型;当`pse_shift`为`bfloat16`类型时,要求`query`为`bfloat16`类型。输入shape类型需为\(B, Q\_N, 1, KV\_S\)或\(1, Q\_N, 1, KV\_S\)。对于`pse_shift`的KV\_S为非32对齐的场景,建议padding到32字节来提高性能,多余部分的填充值不做要求。48+ - Q\_S为1,当`pse_shift`为`float16`类型时,要求`query`为`float16`类型;当`pse_shift`为`bfloat16`类型时,要求`query`为`bfloat16`类型。输入shape类型需为\(B, Q\_N, 1, KV\_S\)或\(1, Q\_N, 1, KV\_S\)。对于`pse_shift`的KV\_S为非32对齐的场景,建议padding到32字节来提高性能,多余部分的填充值不做要求。
48 49 
49-- **atten\_mask**(`Tensor`):可选参数,对QK结果进行mask,用来指示是否计算Token间的相关性。数据类型支持`bool`、`int8`和`uint8`。不支持非连续的Tensor,数据格式支持ND。如不使用该功能可传入None。50+- **atten\_mask**(`Tensor`):可选参数,对QK结果进行mask,用来指示是否计算Token间的相关性。数据类型支持`bool`、`int8`和`uint8`。不支持非连续的Tensor,数据格式支持ND。如不使用该功能可传入None。
50 - `sparse_mode`为0、1时51 - `sparse_mode`为0、1时
51 - 支持shape传入(1,Q_S,KV_S)、(B,1,Q_S,KV_S)、(1,1,Q_S,KV_S)。52 - 支持shape传入(1,Q_S,KV_S)、(B,1,Q_S,KV_S)、(1,1,Q_S,KV_S)。
52 - 当输入`input_layout`为BSH、BSND、BNSD、BNSD_BSND时,且query、key、value的D相等,并且不传`query_rope`和`key_rope`时,Q_S为1可支持传入(B,KV_S),Q_S大于1时可支持传入(Q_S,KV_S)。53 - 当输入`input_layout`为BSH、BSND、BNSD、BNSD_BSND时,且query、key、value的D相等,并且不传`query_rope`和`key_rope`时,Q_S为1可支持传入(B,KV_S),Q_S大于1时可支持传入(Q_S,KV_S)。
@@ -55,111 +56,111 @@ torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=No
55 - `sparse_mode`为9时:56 - `sparse_mode`为9时:
56 - `input_layout`为BSH、BSND或BNSD时,shape输入支持(B, Q_S, Q_S)。57 - `input_layout`为BSH、BSND或BNSD时,shape输入支持(B, Q_S, Q_S)。
57 - `input_layout`为TND时,shape输入支持(∑Q_Si²,),即每个batch的Q_Si×Q_Si mask拼接为1D tensor。58 - `input_layout`为TND时,shape输入支持(∑Q_Si²,),即每个batch的Q_Si×Q_Si mask拼接为1D tensor。
58-- **actual\_seq\_qlen**(`List[Int]`):可选参数,表示不同Batch中`query`的有效seqlen,数据类型支持`int64`。默认值为None,表示和`query`的shape的S长度相同。59+- **actual\_seq\_qlen**(`List[Int]`):可选参数,表示不同Batch中`query`的有效seqlen,数据类型支持`int64`。默认值为None,表示和`query`的shape的S长度相同。
59 该入参中每个Batch的有效seqlen不超过`query`中对应batch的seqlen。当seqlen传入长度为1时,每个Batch使用相同seqlen;当seqlen传入长度>=Batch时,取seqlen的前Batch个数;其他长度不支持。当`query`的input\_layout为TND时,该入参必须传入,且以该入参元素的数量作为Batch值。该入参中每个元素的值表示当前Batch与之前所有Batch的seqlen和,因此后一个元素的值必须>=前一个元素的值,且不能出现负值。60 该入参中每个Batch的有效seqlen不超过`query`中对应batch的seqlen。当seqlen传入长度为1时,每个Batch使用相同seqlen;当seqlen传入长度>=Batch时,取seqlen的前Batch个数;其他长度不支持。当`query`的input\_layout为TND时,该入参必须传入,且以该入参元素的数量作为Batch值。该入参中每个元素的值表示当前Batch与之前所有Batch的seqlen和,因此后一个元素的值必须>=前一个元素的值,且不能出现负值。
60 61 
61-- **actual\_seq\_kvlen**(`List[Int]`):可选参数,表示不同Batch中`key`/`value`的有效seqlenKv,数据类型支持`int64`。默认值为None,表示和key/value的shape的S长度相同。不同O\_S值有不同的约束,具体参见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。62+- **actual\_seq\_kvlen**(`List[Int]`):可选参数,表示不同Batch中`key`/`value`的有效seqlenKv,数据类型支持`int64`。默认值为None,表示和key/value的shape的S长度相同。不同O\_S值有不同的约束,具体参见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
62-- **block\_table**(`Tensor`):可选参数,表示PageAttention中KV存储使用的block映射表,数据类型支持`int32`。数据格式支持ND。如不使用该功能可传入None。63+- **block\_table**(`Tensor`):可选参数,表示PageAttention中KV存储使用的block映射表,数据类型支持`int32`。数据格式支持ND。如不使用该功能可传入None。
63-- **dequant\_scale\_query**(`Tensor`):可选参数,表示`query`的反量化参数,仅支持pertoken叠加perhead。数据类型支持`float32`。数据格式支持ND,如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。64+- **dequant\_scale\_query**(`Tensor`):可选参数,表示`query`的反量化参数,仅支持pertoken叠加perhead。数据类型支持`float32`。数据格式支持ND,如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
64-- **dequant\_scale\_key**(`Tensor`):可选参数,kv伪量化参数分离时表示`key`的反量化因子。数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持ND。通常支持perchannel、pertensor、pertoken、pertensor叠加perhead、pertoken叠加perhead、pertoken叠加使用page attention模式管理scale、pertoken叠加perhead并使用page attention模式管理scale。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。 65+- **dequant\_scale\_key**(`Tensor`):可选参数,kv伪量化参数分离时表示`key`的反量化因子。数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持ND。通常支持perchannel、pertensor、pertoken、pertensor叠加perhead、pertoken叠加perhead、pertoken叠加使用page attention模式管理scale、pertoken叠加perhead并使用page attention模式管理scale。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
65 66
66-- **dequant\_offset\_key**(`Tensor`):可选参数,kv伪量化参数分离时表示`key`的反量化偏移。数据类型支持`float16`、`bfloat16`、`float32`。数据格式支持ND。支持perchannel、pertensor、pertoken、pertensor叠加perhead、pertoken叠加perhead、pertoken叠加使用page attention模式管理offset、pertoken叠加perhead并使用page attention模式管理offset。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。67+- **dequant\_offset\_key**(`Tensor`):可选参数,kv伪量化参数分离时表示`key`的反量化偏移。数据类型支持`float16`、`bfloat16`、`float32`。数据格式支持ND。支持perchannel、pertensor、pertoken、pertensor叠加perhead、pertoken叠加perhead、pertoken叠加使用page attention模式管理offset、pertoken叠加perhead并使用page attention模式管理offset。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
67-- **dequant\_scale\_value**(`Tensor`):可选参数,kv伪量化参数分离时表示`value`的反量化因子。数据类型支持`float16`、`bfloat16`、`float32`。数据格式支持ND。支持perchannel、pertensor、pertoken、pertensor叠加perhead、pertoken叠加perhead、pertoken叠加使用page attention模式管理scale、pertoken叠加perhead并使用page attention模式管理scale。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。68+- **dequant\_scale\_value**(`Tensor`):可选参数,kv伪量化参数分离时表示`value`的反量化因子。数据类型支持`float16`、`bfloat16`、`float32`。数据格式支持ND。支持perchannel、pertensor、pertoken、pertensor叠加perhead、pertoken叠加perhead、pertoken叠加使用page attention模式管理scale、pertoken叠加perhead并使用page attention模式管理scale。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
68 69
69-- **dequant\_offset\_value**(`Tensor`):可选参数,kv伪量化参数分离时表示`value`的反量化偏移。数据类型支持`float16`、`bfloat16`、`float32`。数据格式支持ND。支持perchannel、pertensor、pertoken、pertensor叠加perhead、pertoken叠加perhead、pertoken叠加使用page attention模式管理offset、pertoken叠加perhead并使用page attention模式管理offset。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。70+- **dequant\_offset\_value**(`Tensor`):可选参数,kv伪量化参数分离时表示`value`的反量化偏移。数据类型支持`float16`、`bfloat16`、`float32`。数据格式支持ND。支持perchannel、pertensor、pertoken、pertensor叠加perhead、pertoken叠加perhead、pertoken叠加使用page attention模式管理offset、pertoken叠加perhead并使用page attention模式管理offset。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
70-- **dequant\_scale\_key\_rope**(`Tensor`):可选参数,**预留参数,暂未使用,使用默认值即可。**71+- **dequant\_scale\_key\_rope**(`Tensor`):可选参数,**预留参数,暂未使用,使用默认值即可。**
71-- **quant\_scale\_out**(`Tensor`):可选参数,表示输出的量化因子。数据类型支持`float32`、`bfloat16`。数据格式支持ND。支持pertensor、perchannel。当输入为`bfloat16`时,同时支持`float32`、`bfloat16`,否则仅支持`float32`。perchannel格式,当输出layout为BSH时,要求`quant_scale_out`所有维度的乘积等于H;其他layout要求乘积等于Q\_N\*D(建议输出layout为BSH时,quant\_scale\_out shape传入\(1, 1, H\)或\(H,\);输出为BNSD时,建议传入\(1, Q\_N, 1, D\)或\(Q\_N, D\);输出为BSND时,建议传入\(1, 1, Q\_N, D\)或\(Q\_N, D\))。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。72+- **quant\_scale\_out**(`Tensor`):可选参数,表示输出的量化因子。数据类型支持`float32`、`bfloat16`。数据格式支持ND。支持pertensor、perchannel。当输入为`bfloat16`时,同时支持`float32`、`bfloat16`,否则仅支持`float32`。perchannel格式,当输出layout为BSH时,要求`quant_scale_out`所有维度的乘积等于H;其他layout要求乘积等于Q\_N\*D(建议输出layout为BSH时,quant\_scale\_out shape传入\(1, 1, H\)或\(H,\);输出为BNSD时,建议传入\(1, Q\_N, 1, D\)或\(Q\_N, D\);输出为BSND时,建议传入\(1, 1, Q\_N, D\)或\(Q\_N, D\))。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
72-- **quant\_offset\_out**(`Tensor`):可选参数,表示输出的量化偏移。数据类型支持`float32`、`bfloat16`。数据格式支持ND。支持pertensor、perchannel。若传入`quant_offset_out`,需保证其类型和shape信息与`quant_scale_out`一致。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。73+- **quant\_offset\_out**(`Tensor`):可选参数,表示输出的量化偏移。数据类型支持`float32`、`bfloat16`。数据格式支持ND。支持pertensor、perchannel。若传入`quant_offset_out`,需保证其类型和shape信息与`quant_scale_out`一致。如不使用该功能可传入None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
73-- **learnable_sink**(`Tensor`):可选参数,表示通过可学习的“Sink Token”起到吸收Attention Score的作用,数据类型支持`bfloat16`,数据格式支持ND,shape输入为(Q_N,)。默认值为None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。74+- **learnable_sink**(`Tensor`):可选参数,表示通过可学习的“Sink Token”起到吸收Attention Score的作用,数据类型支持`bfloat16`,数据格式支持ND,shape输入为(Q_N,)。默认值为None,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
74 75 
75-- **num\_query\_heads**(`int`):可选参数,代表query的head个数,数据类型支持`int64`,在BNSD场景下,需要与shape中的`query`的N轴shape值相同,否则执行异常。76+- **num\_query\_heads**(`int`):可选参数,代表query的head个数,数据类型支持`int64`,在BNSD场景下,需要与shape中的`query`的N轴shape值相同,否则执行异常。
76-- **num\_key\_value\_heads**(`int`):可选参数,代表`key`、`value`中head个数,用于支持GQA(Grouped-Query Attention,分组查询注意力)场景,数据类型支持`int64`。默认值为0,表示`key`/`value`/`query`的head个数相等,需要满足`num_query_heads`整除`num_key_value_heads`,`num_query_heads`与`num_key_value_heads`的比值不能大于64。在BSND、BNSD、BNSD\_BSND(仅支持Q\_S大于1)场景下,还需要与shape中的`key`/`value`的N轴shape值相同,否则执行异常。77+- **num\_key\_value\_heads**(`int`):可选参数,代表`key`、`value`中head个数,用于支持GQA(Grouped-Query Attention,分组查询注意力)场景,数据类型支持`int64`。默认值为0,表示`key`/`value`/`query`的head个数相等,需要满足`num_query_heads`整除`num_key_value_heads`,`num_query_heads`与`num_key_value_heads`的比值不能大于64。在BSND、BNSD、BNSD\_BSND(仅支持Q\_S大于1)场景下,还需要与shape中的`key`/`value`的N轴shape值相同,否则执行异常。
77-- **softmax\_scale**(`float`):可选参数,公式中d开根号的倒数,代表缩放系数,作为计算流中Muls的scalar值,数据类型支持`float32`。数据类型与`query`数据类型需满足数据类型推导规则。默认值为1.0。78+- **softmax\_scale**(`float`):可选参数,公式中d开根号的倒数,代表缩放系数,作为计算流中Muls的scalar值,数据类型支持`float32`。数据类型与`query`数据类型需满足数据类型推导规则。默认值为1.0。
78-- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`。默认值为2147483647,Q\_S为1时该参数无效。79+- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`。默认值为2147483647,Q\_S为1时该参数无效。
79-- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`。默认值为2147483647,Q\_S为1时该参数无效。80+- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`。默认值为2147483647,Q\_S为1时该参数无效。
80-- **input\_layout**(`str`):可选参数,用于标识输入`query`、`key`、`value`的数据排布格式,默认值为"BSH"。81+- **input\_layout**(`str`):可选参数,用于标识输入`query`、`key`、`value`的数据排布格式,默认值为"BSH"。
81 82 
82 > [!NOTE] 83 > [!NOTE]
83 > 注意排布格式带下划线时,下划线左边表示输入query的layout,下划线右边表示输出output的格式,算子内部会进行layout转换。84 > 注意排布格式带下划线时,下划线左边表示输入query的layout,下划线右边表示输出output的格式,算子内部会进行layout转换。
84 85 
85 支持BSH、BSND、BNSD、BNSD\_BSND(输入为BNSD时,输出格式为BSND,仅支持Q\_S大于1)、BSH\_NBSD、BSND\_NBSD、BNSD\_NBSD(输出格式为NBSD时,仅支持Q\_S大于1且小于等于16)、TND、TND\_NTD、NTD\_TND(TND相关场景综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214))。其中BNSD\_BSND含义指当输入为BNSD,输出格式为BSND,仅支持Q\_S大于1。86 支持BSH、BSND、BNSD、BNSD\_BSND(输入为BNSD时,输出格式为BSND,仅支持Q\_S大于1)、BSH\_NBSD、BSND\_NBSD、BNSD\_NBSD(输出格式为NBSD时,仅支持Q\_S大于1且小于等于16)、TND、TND\_NTD、NTD\_TND(TND相关场景综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214))。其中BNSD\_BSND含义指当输入为BNSD,输出格式为BSND,仅支持Q\_S大于1。
86 87 
87-- **sparse\_mode**(`int`):可选参数,表示sparse的模式。数据类型支持`int64`。Q\_S为1且不带rope输入时该参数无效。input\_layout为TND、TND\_NTD、NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。88+- **sparse\_mode**(`int`):可选参数,表示sparse的模式。数据类型支持`int64`。Q\_S为1且不带rope输入时该参数无效。input\_layout为TND、TND\_NTD、NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
88 89
89- - `sparse_mode`为0时,代表defaultMask模式,如果atten\_mask未传入则不做mask操作,忽略pre\_tokens和next\_tokens(内部赋值为INT\_MAX);如果传入,则需要传入完整的atten\_mask矩阵(S1\*S2),表示pre\_tokens和next\_tokens之间的部分需要计算。90+ - `sparse_mode`为0时,代表defaultMask模式,如果atten\_mask未传入则不做mask操作,忽略pre\_tokens和next\_tokens(内部赋值为INT\_MAX);如果传入,则需要传入完整的atten\_mask矩阵(S1\*S2),表示pre\_tokens和next\_tokens之间的部分需要计算。
90- - `sparse_mode`为1时,代表allMask,必须传入完整的atten\_mask矩阵(S1\*S2)。91+ - `sparse_mode`为1时,代表allMask,必须传入完整的atten\_mask矩阵(S1\*S2)。
91- - `sparse_mode`为2时,代表leftUpCausal模式的mask,需要传入优化后的atten\_mask矩阵(2048\*2048)。92+ - `sparse_mode`为2时,代表leftUpCausal模式的mask,需要传入优化后的atten\_mask矩阵(2048\*2048)。
92- - `sparse_mode`为3时,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景,需要传入优化后的atten\_mask矩阵(2048\*2048)。93+ - `sparse_mode`为3时,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景,需要传入优化后的atten\_mask矩阵(2048\*2048)。
93- - `sparse_mode`为4时,代表band模式的mask,需要传入优化后的atten\_mask矩阵(2048\*2048)。94+ - `sparse_mode`为4时,代表band模式的mask,需要传入优化后的atten\_mask矩阵(2048\*2048)。
94- - `sparse_mode`为5、6、7、8时,分别代表prefix、global、dilated、block\_local,均暂不支持。95+ - `sparse_mode`为5、6、7、8时,分别代表prefix、global、dilated、block\_local,均暂不支持。
95- - `sparse_mode`为9时,代表treeMask模式,用于推测解码场景的树形注意力掩码。需传入自定义tree mask,仅MLA场景(query\_rope和key\_rope不为空)支持。不支持左padding、pse\_shift、sharedPrefix,输出dtype不支持int8,每个batch需满足Q\_S ≤ KV\_S。默认值为0。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。96+ - `sparse_mode`为9时,代表treeMask模式,用于推测解码场景的树形注意力掩码。需传入自定义tree mask,仅MLA场景(query\_rope和key\_rope不为空)支持。不支持左padding、pse\_shift、sharedPrefix,输出dtype不支持int8,每个batch需满足Q\_S ≤ KV\_S。默认值为0。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
96 97
97-- **block\_size**(`int`):可选参数,表示PageAttention中KV存储每个block中最大的token个数,默认为0,数据类型支持`int64`。98+- **block\_size**(`int`):可选参数,表示PageAttention中KV存储每个block中最大的token个数,默认为0,数据类型支持`int64`。
98-- **query\_quant\_mode**(`int`):可选参数, 表示query的伪量化方式。仅支持传入3,代表模式3:pertoken叠加perhead模式。99+- **query\_quant\_mode**(`int`):可选参数, 表示query的伪量化方式。仅支持传入3,代表模式3:pertoken叠加perhead模式。
99-- **key\_quant\_mode**(`int`):可选参数,表示key的伪量化方式,默认值为0。取值除了`key_quant_mode`为0且`value_quant_mode`为1的场景外,其他场景取值需要与`value_quant_mode`一致。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。100+- **key\_quant\_mode**(`int`):可选参数,表示key的伪量化方式,默认值为0。取值除了`key_quant_mode`为0且`value_quant_mode`为1的场景外,其他场景取值需要与`value_quant_mode`一致。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
100 101 
101 当Q\_S>=2时,仅支持传入值为0、1;当Q\_S=1时,支持取值0、1、2、3、4、5。102 当Q\_S>=2时,仅支持传入值为0、1;当Q\_S=1时,支持取值0、1、2、3、4、5。
102 103 
103- - `key_quant_mode`为0时,代表perchannel模式(perchannel包含pertensor)。104+ - `key_quant_mode`为0时,代表perchannel模式(perchannel包含pertensor)。
104- - `key_quant_mode`为1时,代表pertoken模式。105+ - `key_quant_mode`为1时,代表pertoken模式。
105- - `key_quant_mode`为2时,代表pertensor叠加perhead模式。106+ - `key_quant_mode`为2时,代表pertensor叠加perhead模式。
106- - `key_quant_mode`为3时,代表pertoken叠加perhead模式。107+ - `key_quant_mode`为3时,代表pertoken叠加perhead模式。
107- - `key_quant_mode`为4时,代表pertoken叠加使用page attention模式管理scale/offset模式。108+ - `key_quant_mode`为4时,代表pertoken叠加使用page attention模式管理scale/offset模式。
108- - `key_quant_mode`为5时,代表pertoken叠加perhead并使用page attention模式管理scale/offset模式。109+ - `key_quant_mode`为5时,代表pertoken叠加perhead并使用page attention模式管理scale/offset模式。
109 110 
110- **value\_quant\_mode**`int`):可选参数,表示`value`的伪量化方式,模式编号与`key_quant_mode`一致,默认值为0。取值除了`key_quant_mode`为0且`value_quant_mode`为1的场景外,其他场景取值需要与`key_quant_mode`一致。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。111- **value\_quant\_mode**`int`):可选参数,表示`value`的伪量化方式,模式编号与`key_quant_mode`一致,默认值为0。取值除了`key_quant_mode`为0且`value_quant_mode`为1的场景外,其他场景取值需要与`key_quant_mode`一致。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
111 112 
112 当Q\_S>=2时,仅支持传入值为0、1;当Q\_S=1时,支持取值0、1、2、3、4、5。113 当Q\_S>=2时,仅支持传入值为0、1;当Q\_S=1时,支持取值0、1、2、3、4、5。
113 114 
114-- **inner\_precise**(`int`):可选参数,数据类型支持`int64`,支持4种模式:0、1、2、3。一共两位bit位,第0位(bit0)表示高精度或者高性能选择,第1位(bit1)表示是否做行无效修正。当Q\_S\>1时,sparse\_mode为0或1,并传入用户自定义mask的情况下,建议开启行无效;Q\_S为1时该参数仅支持取0和1。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。115+- **inner\_precise**(`int`):可选参数,数据类型支持`int64`,支持4种模式:0、1、2、3。一共两位bit位,第0位(bit0)表示高精度或者高性能选择,第1位(bit1)表示是否做行无效修正。当Q\_S\>1时,sparse\_mode为0或1,并传入用户自定义mask的情况下,建议开启行无效;Q\_S为1时该参数仅支持取0和1。综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
115 116 
116- - inner\_precise为0时,代表开启高精度模式,且不做行无效修正。117+ - inner\_precise为0时,代表开启高精度模式,且不做行无效修正。
117- - inner\_precise为1时,代表高性能模式,且不做行无效修正。118+ - inner\_precise为1时,代表高性能模式,且不做行无效修正。
118- - inner\_precise为2时,代表开启高精度模式,且做行无效修正。119+ - inner\_precise为2时,代表开启高精度模式,且做行无效修正。
119- - inner\_precise为3时,代表高性能模式,且做行无效修正。120+ - inner\_precise为3时,代表高性能模式,且做行无效修正。
120 121 
121 > [!NOTE] 122 > [!NOTE]
122 > bfloat16和int8不区分高精度和高性能,行无效修正对`float16`、`bfloat16`和`int8`均生效。当前0、1为保留配置值,当计算过程中“参与计算的mask部分”存在某整行全为1的情况时,精度可能会有损失。此时可以尝试将该参数配置为2或3来使能行无效功能以提升精度,但是该配置会导致性能下降。123 > bfloat16和int8不区分高精度和高性能,行无效修正对`float16`、`bfloat16`和`int8`均生效。当前0、1为保留配置值,当计算过程中“参与计算的mask部分”存在某整行全为1的情况时,精度可能会有损失。此时可以尝试将该参数配置为2或3来使能行无效功能以提升精度,但是该配置会导致性能下降。
123 124 
124-- **return\_softmax\_lse**(`bool`):可选参数,表示是否输出`softmax_lse`,支持S轴外切(增加输出)。true表示输出,false表示不输出;默认值为false。125+- **return\_softmax\_lse**(`bool`):可选参数,表示是否输出`softmax_lse`,支持S轴外切(增加输出)。true表示输出,false表示不输出;默认值为false。
125-- **query_dtype**(`int`):可选参数,表示`query`的数据类型,**预留参数,暂未使用,使用默认值即可。**126+- **query_dtype**(`int`):可选参数,表示`query`的数据类型,**预留参数,暂未使用,使用默认值即可。**
126-- **key_dtype**(`int`):可选参数,表示`key`的数据类型,**预留参数,暂未使用,使用默认值即可。**127+- **key_dtype**(`int`):可选参数,表示`key`的数据类型,**预留参数,暂未使用,使用默认值即可。**
127-- **value_dtype**(`int`):可选参数,表示`value`的数据类型,**预留参数,暂未使用,使用默认值即可。**128+- **value_dtype**(`int`):可选参数,表示`value`的数据类型,**预留参数,暂未使用,使用默认值即可。**
128-- **query_rope_dtype**(`int`):可选参数,表示`query_rope`的数据类型,**预留参数,暂未使用,使用默认值即可。**129+- **query_rope_dtype**(`int`):可选参数,表示`query_rope`的数据类型,**预留参数,暂未使用,使用默认值即可。**
129-- **key_rope_dtype**(`int`):可选参数,表示`key_rope`的数据类型,**预留参数,暂未使用,使用默认值即可。**130+- **key_rope_dtype**(`int`):可选参数,表示`key_rope`的数据类型,**预留参数,暂未使用,使用默认值即可。**
130-- **key_shared_prefix_dtype**(`int`):可选参数,表示key_shared_prefix的数据类型,**预留参数,暂未使用,使用默认值即可。**131+- **key_shared_prefix_dtype**(`int`):可选参数,表示key_shared_prefix的数据类型,**预留参数,暂未使用,使用默认值即可。**
131-- **value_shared_prefix_dtype**(`int`):可选参数,表示value_shared_prefix的数据类型,**预留参数,暂未使用,使用默认值即可。**132+- **value_shared_prefix_dtype**(`int`):可选参数,表示value_shared_prefix的数据类型,**预留参数,暂未使用,使用默认值即可。**
132-- **dequant_scale_query_dtype**(`int`):可选参数,表示`dequant_scale_query`的数据类型,**预留参数,暂未使用,使用默认值即可。**133+- **dequant_scale_query_dtype**(`int`):可选参数,表示`dequant_scale_query`的数据类型,**预留参数,暂未使用,使用默认值即可。**
133-- **dequant_scale_key_dtype**(`int`):可选参数,表示`dequant_scale_key`的数据类型,**预留参数,暂未使用,使用默认值即可。**134+- **dequant_scale_key_dtype**(`int`):可选参数,表示`dequant_scale_key`的数据类型,**预留参数,暂未使用,使用默认值即可。**
134-- **dequant_scale_value_dtype**(`int`):可选参数,表示`dequant_scale_value`的数据类型,**预留参数,暂未使用,使用默认值即可。**135+- **dequant_scale_value_dtype**(`int`):可选参数,表示`dequant_scale_value`的数据类型,**预留参数,暂未使用,使用默认值即可。**
135-- **dequant_scale_key_rope_dtype**(`int`):可选参数,表示`dequant_scale_key_rope`的数据类型,**预留参数,暂未使用,使用默认值即可。**136+- **dequant_scale_key_rope_dtype**(`int`):可选参数,表示`dequant_scale_key_rope`的数据类型,**预留参数,暂未使用,使用默认值即可。**
136 137 
137## 返回值说明<a name="zh-cn_topic_0000001832267082_section22231435517"></a>138## 返回值说明<a name="zh-cn_topic_0000001832267082_section22231435517"></a>
138 139 
139-- **attention\_out**(`Tensor`):公式中的输出,数据类型支持`float16`、`bfloat16`、`int8`。数据格式支持ND。限制:该入参的D维度与`value`的D保持一致,其余维度需要与入参`query`的shape保持一致。140+- **attention\_out**(`Tensor`):公式中的输出,数据类型支持`float16`、`bfloat16`、`int8`。数据格式支持ND。限制:该入参的D维度与`value`的D保持一致,其余维度需要与入参`query`的shape保持一致。
140-- **softmax\_lse**(`Tensor`):ring attention算法对query乘key的结果先取max得到softmax\_max,query乘key的结果减去softmax\_max,再取exp,最后取sum,得到softmax\_sum,最后对softmax\_sum取log,再加上softmax\_max得到的结果。数据类型支持`float32`,当`return_softmax_lse`为True时,一般情况下输出shape为\(B, Q\_N, Q\_S, 1\),若input\_layout为TND/NTD\_TND时,输出shape为\(T,Q\_N,1\);当`return_softmax_lse`为False时,输出shape为\[1\]的值为0的Tensor。141+- **softmax\_lse**(`Tensor`):ring attention算法对query乘key的结果先取max得到softmax\_max,query乘key的结果减去softmax\_max,再取exp,最后取sum,得到softmax\_sum,最后对softmax\_sum取log,再加上softmax\_max得到的结果。数据类型支持`float32`,当`return_softmax_lse`为True时,一般情况下输出shape为\(B, Q\_N, Q\_S, 1\),若input\_layout为TND/NTD\_TND时,输出shape为\(T,Q\_N,1\);当`return_softmax_lse`为False时,输出shape为\[1\]的值为0的Tensor。
141 142 
142## 约束说明<a name="zh-cn_topic_0000001832267082_section12345537164214"></a>143## 约束说明<a name="zh-cn_topic_0000001832267082_section12345537164214"></a>
143 144 
144-- 该接口支持推理场景下使用。145+- 该接口支持推理场景下使用。
145-- 该接口支持图模式。146+- 该接口支持图模式。
146-- 该接口与PyTorch配合使用时,需要保证CANN相关包与PyTorch相关包的版本匹配。147+- 该接口与PyTorch配合使用时,需要保证CANN相关包与PyTorch相关包的版本匹配。
147-- 入参为空的处理:算子内部需要判断参数query是否为空,如果是空则直接返回空。参数query不为空Tensor,参数key、value为空Tensor(即S2为0),则attention\_out按照对应shape大小返回全0。attention\_out为空Tensor时,返回空。148+- 入参为空的处理:算子内部需要判断参数query是否为空,如果是空则直接返回空。参数query不为空Tensor,参数key、value为空Tensor(即S2为0),则attention\_out按照对应shape大小返回全0。attention\_out为空Tensor时,返回空。
148-- 参数key、value中对应tensor的shape需要完全一致;非连续场景下key、value的tensorlist中的batch只能为1,个数等于query的B,N和D需要相等。149+- 参数key、value中对应tensor的shape需要完全一致;非连续场景下key、value的tensorlist中的batch只能为1,个数等于query的B,N和D需要相等。
149-- int8量化相关入参数量与输出数据格式的综合限制:150+- int8量化相关入参数量与输出数据格式的综合限制:
150- - 输出为int8的场景:入参quant\_scale\_out需要存在,quant\_offset\_out可选,不传时默认为0。151+ - 输出为int8的场景:入参quant\_scale\_out需要存在,quant\_offset\_out可选,不传时默认为0。
151- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:输入为int8。152+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:输入为int8。
152- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:输入为int8。153+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:输入为int8。
153 154 
154- - 输出为float16的场景:若存在入参quant\_offset\_out或quant\_scale\_out(即不为None),则报错并返回。155+ - 输出为float16的场景:若存在入参quant\_offset\_out或quant\_scale\_out(即不为None),则报错并返回。
155- - 入参quant\_offset\_out和quant\_scale\_out支持pertensor或perchannel格式,数据类型支持float32、bfloat16。156+ - 入参quant\_offset\_out和quant\_scale\_out支持pertensor或perchannel格式,数据类型支持float32、bfloat16。
156 157 
157-- query\_rope和key\_rope输入时即为MLA场景,参数约束如下:158+- query\_rope和key\_rope输入时即为MLA场景,参数约束如下:
158- - query\_rope的数据类型、数据格式与query一致。159+ - query\_rope的数据类型、数据格式与query一致。
159- - key\_rope的数据类型、数据格式与key一致。160+ - key\_rope的数据类型、数据格式与key一致。
160- - query\_rope和key\_rope要求同时配置或同时不配置,不支持只配置其中一个。161+ - query\_rope和key\_rope要求同时配置或同时不配置,不支持只配置其中一个。
161- - 当query\_rope和key\_rope非空时,query的D只支持512、128;162+ - 当query\_rope和key\_rope非空时,query的D只支持512、128;
162- - 当query的D等于512时:163+ - 当query的D等于512时:
163 - sparse:支持0/3/4/9;164 - sparse:支持0/3/4/9;
164 - query\_rope配置时要求query的N为1/2/4/8/16/32/64/128,query\_rope的shape中D为64,其余维度与query一致;165 - query\_rope配置时要求query的N为1/2/4/8/16/32/64/128,query\_rope的shape中D为64,其余维度与query一致;
165 - key\_rope配置时要求key的N为1、D为512,key\_rope的shape中D为64,其余维度与key一致;166 - key\_rope配置时要求key的N为1、D为512,key\_rope的shape中D为64,其余维度与key一致;
@@ -172,30 +173,30 @@ torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=No
172 - 不支持传入quant\_scale\_out、quant\_offset\_out、dequant\_offset\_key、dequant\_offset\_value,否则报错并返回。173 - 不支持传入quant\_scale\_out、quant\_offset\_out、dequant\_offset\_key、dequant\_offset\_value,否则报错并返回。
173 - query\_quant\_mode仅支持pertoken叠加perhead模式,key\_quant\_mode和value\_quant\_mode仅支持pertensor模式。174 - query\_quant\_mode仅支持pertoken叠加perhead模式,key\_quant\_mode和value\_quant\_mode仅支持pertensor模式。
174 - 支持key、value、key\_rope的input\_layout格式为NZ。175 - 支持key、value、key\_rope的input\_layout格式为NZ。
175- - 当query的D等于128时:176+ - 当query的D等于128时:
176 - input\_layout:BSH、BSND、TND、BNSD、NTD、BSH\_BNSD、BSND\_BNSD、BNSD\_BSND、NTD\_TND。 177 - input\_layout:BSH、BSND、TND、BNSD、NTD、BSH\_BNSD、BSND\_BNSD、BNSD\_BSND、NTD\_TND。
177 - query\_rope配置时要求query\_rope的shape中D为64,其余维度与query一致。 178 - query\_rope配置时要求query\_rope的shape中D为64,其余维度与query一致。
178 - key\_rope配置时要求key\_rope的shape中D为64,其余维度与key一致。 179 - key\_rope配置时要求key\_rope的shape中D为64,其余维度与key一致。
179 - 不支持开启左padding、tensorlist、pse、prefix、伪量化、全量化、后量化、空Tensor。180 - 不支持开启左padding、tensorlist、pse、prefix、伪量化、全量化、后量化、空Tensor。
180 - 其余约束同TND、NTD\_TND场景下的综合限制保持一致。181 - 其余约束同TND、NTD\_TND场景下的综合限制保持一致。
181 182 
182- - TND、TND\_NTD、NTD\_TND场景下query、key、value输入的综合限制:183+ - TND、TND\_NTD、NTD\_TND场景下query、key、value输入的综合限制:
183- - actual\_seq\_qlen和actual\_seq\_kvlen必须传入,且以该入参元素数量作为Batch值(注意入参元素数量要小于等于4096)。该入参中每个元素的值表示当前Batch与之前所有Batch的Sequence Length和,因此后一个元素的值必须大于等于前一个元素的值;184+ - actual\_seq\_qlen和actual\_seq\_kvlen必须传入,且以该入参元素数量作为Batch值(注意入参元素数量要小于等于4096)。该入参中每个元素的值表示当前Batch与之前所有Batch的Sequence Length和,因此后一个元素的值必须大于等于前一个元素的值;
184- - 当query的D等于512时:185+ - 当query的D等于512时:
185- - sparse:支持0/3/4/9;186+ - sparse:支持0/3/4/9;
186- - 支持TND、TND\_NTD;187+ - 支持TND、TND\_NTD;
187- - 支持开启page attention,此时actual\_seq\_kvlen长度等于key/value的batch值,代表每个batch的实际长度,值不大于KV\_S;188+ - 支持开启page attention,此时actual\_seq\_kvlen长度等于key/value的batch值,代表每个batch的实际长度,值不大于KV\_S;
188- - 要求query的N为1/2/4/8/16/32/64/128,key、value的N为1;189+ - 要求query的N为1/2/4/8/16/32/64/128,key、value的N为1;
189- - 要求query\_rope和key\_rope不等于空,query\_rope和key\_rope的D为64;190+ - 要求query\_rope和key\_rope不等于空,query\_rope和key\_rope的D为64;
190- - 不支持开启左padding、tensorlist、pse、prefix、伪量化、全量化、后量化、空Tensor。191+ - 不支持开启左padding、tensorlist、pse、prefix、伪量化、全量化、后量化、空Tensor。
191 192 
192- - 当query的D不等于512时:193+ - 当query的D不等于512时:
193- - 当query\_rope和key\_rope为空时:TND场景,要求Q\_D(`query`的D维度)、K\_D(`key`的D维度)、V\_D(`value`的D维度)等于128,或者Q\_D、K\_D等于192,V\_D等于128/192;NTD场景,不支持V\_D等于192;NTD\_TND场景,要求Q\_D、K\_D等于128/192,V\_D等于128。当query\_rope和key\_rope不为空时,要求Q\_D、K\_D、V\_D等于128;GQA和PA场景不支持V_D等于192; MHA(Multi-Head Attention)场景Q\_D、K\_D、V\_D都等于64,或Q\_D、K\_D、V\_D都等于128,或Q\_D和K\_D等于192时V\_D等于128,194+ - 当query\_rope和key\_rope为空时:TND场景,要求Q\_D(`query`的D维度)、K\_D(`key`的D维度)、V\_D(`value`的D维度)等于128,或者Q\_D、K\_D等于192,V\_D等于128/192;NTD场景,不支持V\_D等于192;NTD\_TND场景,要求Q\_D、K\_D等于128/192,V\_D等于128。当query\_rope和key\_rope不为空时,要求Q\_D、K\_D、V\_D等于128;GQA和PA场景不支持V_D等于192; MHA(Multi-Head Attention)场景Q\_D、K\_D、V\_D都等于64,或Q\_D、K\_D、V\_D都等于128,或Q\_D和K\_D等于192时V\_D等于128,
194- - 支持TND、NTD、NTD\_TND;195+ - 支持TND、NTD、NTD\_TND;
195- - page attention场景下仅支持blocksize为16对齐且小于等于1024;196+ - page attention场景下仅支持blocksize为16对齐且小于等于1024;
196- - MHA场景下仅支持数据类型为`float16`、`bfloat16`,当数据类型为`float16`,inner\_precise仅支持0和1,当数据类型为`bfloat16`,inner\_precise仅支持0。当sparse\_mode=0不传atten\_mask矩阵,sparse\_mode为3/4传优化后的atten\_mask矩阵。page attention仅支持BnBsH格式,BnBsH表示KV Cache的排布格式为(blockNum, blocksize, H),其中blockNum为块数量、blocksize为每个块中的token个数、H为隐藏层大小;197+ - MHA场景下仅支持数据类型为`float16`、`bfloat16`,当数据类型为`float16`,inner\_precise仅支持0和1,当数据类型为`bfloat16`,inner\_precise仅支持0。当sparse\_mode=0不传atten\_mask矩阵,sparse\_mode为3/4传优化后的atten\_mask矩阵。page attention仅支持BnBsH格式,BnBsH表示KV Cache的排布格式为(blockNum, blocksize, H),其中blockNum为块数量、blocksize为每个块中的token个数、H为隐藏层大小;
197- - 不支持开启左padding、tensorlist、pse、prefix、伪量化、全量化;198+ - 不支持开启左padding、tensorlist、pse、prefix、伪量化、全量化;
198-- GQA伪量化场景下KV为NZ格式时的参数约束如下:199+- GQA伪量化场景下KV为NZ格式时的参数约束如下:
199 - 支持perchannel和pertoken模式,query数据类型固定为bfloat16,key&value固定为int8;query&key&value的D仅支持128;query Sequence Length仅支持1-16;200 - 支持perchannel和pertoken模式,query数据类型固定为bfloat16,key&value固定为int8;query&key&value的D仅支持128;query Sequence Length仅支持1-16;
200 - input\_layout仅支持BSH、BSND、BNSD;201 - input\_layout仅支持BSH、BSND、BNSD;
201 - 仅支持page_attention场景,blockSize仅支持128或512;202 - 仅支持page_attention场景,blockSize仅支持128或512;
@@ -209,78 +210,78 @@ torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=No
209 - 不支持配置query\_rope和key\_rope;210 - 不支持配置query\_rope和key\_rope;
210 - 不支持左padding、tensorlist、pse、prefix、后量化;211 - 不支持左padding、tensorlist、pse、prefix、后量化;
211 - num\_query\_heads与num\_key\_value\_heads支持组合有(10, 1)、(64, 8)、(80, 8)、(128, 16)。212 - num\_query\_heads与num\_key\_value\_heads支持组合有(10, 1)、(64, 8)、(80, 8)、(128, 16)。
212-- learnable_sink的参数约束如下:213+- learnable_sink的参数约束如下:
213 - 仅支持TND、NTD\_TND;214 - 仅支持TND、NTD\_TND;
214 - 仅支持value的D小于等于128;215 - 仅支持value的D小于等于128;
215 - 仅支持非量化场景。216 - 仅支持非量化场景。
216 - 不支持pse、左padding、公共前缀、后量化。217 - 不支持pse、左padding、公共前缀、后量化。
217-- **当Q\_S大于1时:**218+- **当Q\_S大于1时:**
218- - query、key、value输入,功能使用限制如下:219+ - query、key、value输入,功能使用限制如下:
219- - 支持B轴小于等于65536,D轴32byte不对齐时仅支持到128。220+ - 支持B轴小于等于65536,D轴32byte不对齐时仅支持到128。
220- - 支持N轴小于等于256,支持D轴小于等于512;input\_layout为BSH或者BSND时,建议N\*D小于65535。221+ - 支持N轴小于等于256,支持D轴小于等于512;input\_layout为BSH或者BSND时,建议N\*D小于65535。
221- - S支持小于等于20971520(20M)。部分长序列场景下,如果计算量过大可能会导致PFA算子执行超时(aicore error类型报错,errorStr为timeout or trap error),此场景下建议做S切分处理(注:这里计算量会受B、S、N、D等的影响,值越大计算量越大),典型的会超时的长序列(即B、S、N、D的乘积较大)场景包括但不限于:222+ - S支持小于等于20971520(20M)。部分长序列场景下,如果计算量过大可能会导致PFA算子执行超时(aicore error类型报错,errorStr为timeout or trap error),此场景下建议做S切分处理(注:这里计算量会受B、S、N、D等的影响,值越大计算量越大),典型的会超时的长序列(即B、S、N、D的乘积较大)场景包括但不限于:
222- - B=1,Q\_N=20,Q\_S=2097152,D=256,KV\_N=1,KV\_S=2097152。223+ - B=1,Q\_N=20,Q\_S=2097152,D=256,KV\_N=1,KV\_S=2097152。
223- - B=1,Q\_N=2,Q\_S=20971520,D=256,KV\_N=2,KV\_S=20971520。224+ - B=1,Q\_N=2,Q\_S=20971520,D=256,KV\_N=2,KV\_S=20971520。
224- - B=20,Q\_N=1,Q\_S=2097152,D=256,KV\_N=1,KV\_S=2097152。225+ - B=20,Q\_N=1,Q\_S=2097152,D=256,KV\_N=1,KV\_S=2097152。
225- - B=1,Q\_N=10,Q\_S=2097152,D=512,KV\_N=1,KV\_S=2097152。226+ - B=1,Q\_N=10,Q\_S=2097152,D=512,KV\_N=1,KV\_S=2097152。
226 227 
227- - query、key、value输入类型包含int8时,D轴需要32对齐;输入类型全为`float16`、`bfloat16`时,D轴需16对齐。228+ - query、key、value输入类型包含int8时,D轴需要32对齐;输入类型全为`float16`、`bfloat16`时,D轴需16对齐。
228- - D轴限制:229+ - D轴限制:
229- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:query、key、value输入类型包含`int8`时,D轴需要32对齐;query、key、value或attentionOut类型包含`int4`时,D轴需要64对齐;输入类型全为`float16`、`bfloat16`时,D轴需16对齐。230+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:query、key、value输入类型包含`int8`时,D轴需要32对齐;query、key、value或attentionOut类型包含`int4`时,D轴需要64对齐;输入类型全为`float16`、`bfloat16`时,D轴需16对齐。
230 231 
231- - actual\_seq\_qlen:232+ - actual\_seq\_qlen:
232 233
233 <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于query中对应batch的Sequence Length。seqlen的传入长度为1时,每个Batch使用相同seqlen;传入长度大于等于Batch时取seqlen的前Batch个数。其他长度不支持。当query的input\_layout为TND/NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。234 <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于query中对应batch的Sequence Length。seqlen的传入长度为1时,每个Batch使用相同seqlen;传入长度大于等于Batch时取seqlen的前Batch个数。其他长度不支持。当query的input\_layout为TND/NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
234 235
235- - actual\_seq\_kvlen:236+ - actual\_seq\_kvlen:
236 237
237 <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于key/value中对应batch的Sequence Length。seqlenKv的传入长度为1时,每个Batch使用相同seqlenKv;传入长度大于等于Batch时取seqlenKv的前Batch个数。其他长度不支持。当key/value的input\_layout为TND/NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。238 <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于key/value中对应batch的Sequence Length。seqlenKv的传入长度为1时,每个Batch使用相同seqlenKv;传入长度大于等于Batch时取seqlenKv的前Batch个数。其他长度不支持。当key/value的input\_layout为TND/NTD\_TND时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
238 239
239- - 参数sparse\_mode当前仅支持值为0、1、2、3、4、9的场景,取其它值时会报错。240+ - 参数sparse\_mode当前仅支持值为0、1、2、3、4、9的场景,取其它值时会报错。
240 241 
241- - sparse\_mode=0时,atten\_mask如果为None,则忽略入参pre\_tokens、next\_tokens(内部赋值为INT\_MAX)。242+ - sparse\_mode=0时,atten\_mask如果为None,则忽略入参pre\_tokens、next\_tokens(内部赋值为INT\_MAX)。
242- - sparse\_mode=2、3、4时,atten\_mask的shape需要为\(S, S\)或\(1, S, S\)或\(1, 1, S, S\),其中S的值需要固定为2048,且需要用户保证传入的atten\_mask为下三角,不传入atten\_mask或者传入的shape不正确报错。243+ - sparse\_mode=2、3、4时,atten\_mask的shape需要为\(S, S\)或\(1, S, S\)或\(1, 1, S, S\),其中S的值需要固定为2048,且需要用户保证传入的atten\_mask为下三角,不传入atten\_mask或者传入的shape不正确报错。
243- - sparse\_mode=1、2、3的场景忽略入参pre\_tokens、next\_tokens并按照相关规则赋值。244+ - sparse\_mode=1、2、3的场景忽略入参pre\_tokens、next\_tokens并按照相关规则赋值。
244- - sparse\_mode=9时,仅MLA场景(query\_rope和key\_rope不为空)支持。atten\_mask不能为None。input\_layout为BSH/BSND/BNSD时shape为\(B, Q\_S, Q\_S\);input\_layout为TND时shape为\(∑Q\_Si², \)。不支持左padding、pse\_shift、sharedPrefix。245+ - sparse\_mode=9时,仅MLA场景(query\_rope和key\_rope不为空)支持。atten\_mask不能为None。input\_layout为BSH/BSND/BNSD时shape为\(B, Q\_S, Q\_S\);input\_layout为TND时shape为\(∑Q\_Si², \)。不支持左padding、pse\_shift、sharedPrefix。
245 246 
246- - page attention场景:247+ - page attention场景:
247- - page attention的使能必要条件是block\_table存在且有效,同时key、value是按照block\_table中的索引在一片连续内存中排布,支持key、value数据类型为`float16`、`bfloat16`。在该场景下key、value的input\_layout参数无效。block\_table中填充的是blockid,当前不会对blockid的合法性进行校验,需用户自行保证。248+ - page attention的使能必要条件是block\_table存在且有效,同时key、value是按照block\_table中的索引在一片连续内存中排布,支持key、value数据类型为`float16`、`bfloat16`。在该场景下key、value的input\_layout参数无效。block\_table中填充的是blockid,当前不会对blockid的合法性进行校验,需用户自行保证。
248- - block\_size是用户自定义的参数,该参数的取值会影响page attention的性能,在使能page attention场景下,block\_size最小为128,最大为512,且要求是128的倍数。通常情况下,page attention可以提高吞吐量,但会带来性能上的下降。249+ - block\_size是用户自定义的参数,该参数的取值会影响page attention的性能,在使能page attention场景下,block\_size最小为128,最大为512,且要求是128的倍数。通常情况下,page attention可以提高吞吐量,但会带来性能上的下降。
249 250
250- - page attention场景下,当输入kv cache排布格式为(blocknum, blocksize, H),且KV\_N\*D超过65535时,受硬件指令约束,会被拦截报错。可通过使能GQA(减小KV\_N)或调整kv cache排布格式为(blocknum, KV\_N, blocksize, D)解决。当query的input\_layout为BNSD、TND时,kv cache排布支持(blocknum, blocksize, H)和(blocknum, KV\_N, blocksize, D)两种格式,当query的input\_layout为BSH、BSND时,kv cache排布只支持(blocknum, blocksize, H)一种格式。blocknum不能小于根据actual\_seq\_kvlen和blockSize计算的每个batch的block数量之和。且key和value的shape需保证一致。251+ - page attention场景下,当输入kv cache排布格式为(blocknum, blocksize, H),且KV\_N\*D超过65535时,受硬件指令约束,会被拦截报错。可通过使能GQA(减小KV\_N)或调整kv cache排布格式为(blocknum, KV\_N, blocksize, D)解决。当query的input\_layout为BNSD、TND时,kv cache排布支持(blocknum, blocksize, H)和(blocknum, KV\_N, blocksize, D)两种格式,当query的input\_layout为BSH、BSND时,kv cache排布只支持(blocknum, blocksize, H)一种格式。blocknum不能小于根据actual\_seq\_kvlen和blockSize计算的每个batch的block数量之和。且key和value的shape需保证一致。
251- - page attention不支持伪量化场景,不支持tensorlist场景。252+ - page attention不支持伪量化场景,不支持tensorlist场景。
252- - page attention场景下,必须传入actual\_seq\_kvlen。253+ - page attention场景下,必须传入actual\_seq\_kvlen。
253- - page attention场景下,block\_table必须为二维,第一维长度需等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为不同batch中最大actual\_seq\_kvlen对应的block数量)。254+ - page attention场景下,block\_table必须为二维,第一维长度需等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为不同batch中最大actual\_seq\_kvlen对应的block数量)。
254- - page attention场景下,支持两种格式和float32/bfloat16,不支持输入query为int8的场景。255+ - page attention场景下,支持两种格式和float32/bfloat16,不支持输入query为int8的场景。
255- - page attention使能场景下,以下场景输入需满足KV\_S\>=maxBlockNumPerSeq\*blockSize:256+ - page attention使能场景下,以下场景输入需满足KV\_S\>=maxBlockNumPerSeq\*blockSize:
256- - 传入atten\_mask时,如mask shape为(B, 1, Q\_S, KV\_S)。257+ - 传入atten\_mask时,如mask shape为(B, 1, Q\_S, KV\_S)。
257- - 传入pse\_shift时,如pse\_shift shape为(B, Q\_N, Q\_S, KV\_S)。258+ - 传入pse\_shift时,如pse\_shift shape为(B, Q\_N, Q\_S, KV\_S)。
258 259 
259- - 入参quant\_scale\_out和quant\_offset\_out支持pertensor、perchannel量化,支持float32、bfloat16类型。若传入quant\_offset\_out,需保证其类型和shape信息与quant\_scale\_out一致。当输入为bfloat16时,同时支持float32和bfloat16,否则仅支持float32。perchannel场景下,当输出layout为BSH时,要求quant\_scale\_out所有维度的乘积等于H;其他layout要求乘积等于Q\_N\*D。当输出layout为BSH时,quant\_scale\_out shape建议传入\(1, 1, H\)或\(H,\);当输出layout为BNSD时,建议传入\(1, Q\_N, 1, D\)或\(Q\_N, D\);当输出为BSND时,建议传入\(1, 1, Q\_N, D\)或\(Q\_N, D)。260+ - 入参quant\_scale\_out和quant\_offset\_out支持pertensor、perchannel量化,支持float32、bfloat16类型。若传入quant\_offset\_out,需保证其类型和shape信息与quant\_scale\_out一致。当输入为bfloat16时,同时支持float32和bfloat16,否则仅支持float32。perchannel场景下,当输出layout为BSH时,要求quant\_scale\_out所有维度的乘积等于H;其他layout要求乘积等于Q\_N\*D。当输出layout为BSH时,quant\_scale\_out shape建议传入\(1, 1, H\)或\(H,\);当输出layout为BNSD时,建议传入\(1, Q\_N, 1, D\)或\(Q\_N, D\);当输出为BSND时,建议传入\(1, 1, Q\_N, D\)或\(Q\_N, D)。
260- - 输出为int8,quant\_scale\_out和quant\_offset\_out为perchannel时,暂不支持Ring Attention或者D非32Byte对齐的场景。261+ - 输出为int8,quant\_scale\_out和quant\_offset\_out为perchannel时,暂不支持Ring Attention或者D非32Byte对齐的场景。
261- - 输出为int8时,暂不支持sparse为band且preTokens/nextTokens为负数。262+ - 输出为int8时,暂不支持sparse为band且preTokens/nextTokens为负数。
262- - pse\_shift功能使用限制如下:263+ - pse\_shift功能使用限制如下:
263 264
264- - 支持query数据类型为float16、bfloat16、int8场景下使用该功能。265+ - 支持query数据类型为float16、bfloat16、int8场景下使用该功能。
265- - query、key、value数据类型为float16且pse\_shift存在时,强制走高精度模式,对应的限制继承自高精度模式的限制。266+ - query、key、value数据类型为float16且pse\_shift存在时,强制走高精度模式,对应的限制继承自高精度模式的限制。
266- - Q\_S需大于等于query的S长度,KV\_S需大于等于key的S长度。267+ - Q\_S需大于等于query的S长度,KV\_S需大于等于key的S长度。
267 268 
268- - 输出为int8,入参quant\_offset\_out传入非None和非空tensor值,并且sparse\_mode、pre\_tokens和next\_tokens满足以下条件,矩阵会存在某几行不参与计算的情况,导致计算结果误差,该场景会拦截:269+ - 输出为int8,入参quant\_offset\_out传入非None和非空tensor值,并且sparse\_mode、pre\_tokens和next\_tokens满足以下条件,矩阵会存在某几行不参与计算的情况,导致计算结果误差,该场景会拦截:
269- - sparse\_mode=0,atten\_mask如果非None,每个batch actual\_seq\_qlen-actual\_seq\_kvlen-pre\_tokens\>0或next\_tokens<0时,满足拦截条件。270+ - sparse\_mode=0,atten\_mask如果非None,每个batch actual\_seq\_qlen-actual\_seq\_kvlen-pre\_tokens\>0或next\_tokens<0时,满足拦截条件。
270- - sparse\_mode=1或2,不会出现满足拦截条件的情况。271+ - sparse\_mode=1或2,不会出现满足拦截条件的情况。
271- - sparse\_mode=3,每个batch actual\_seq\_kvlen-actual\_seq\_qlen<0,满足拦截条件。272+ - sparse\_mode=3,每个batch actual\_seq\_kvlen-actual\_seq\_qlen<0,满足拦截条件。
272- - sparse\_mode=4,pre\_tokens<0或每个batch next\_tokens+actual\_seq\_kvlen-actual\_seq\_qlen<0时,满足拦截条件。273+ - sparse\_mode=4,pre\_tokens<0或每个batch next\_tokens+actual\_seq\_kvlen-actual\_seq\_qlen<0时,满足拦截条件。
273 274 
274- - kv伪量化参数分离:275+ - kv伪量化参数分离:
275- - 当伪量化参数和KV分离量化参数同时传入时,以KV分离量化参数为准。 276+ - 当伪量化参数和KV分离量化参数同时传入时,以KV分离量化参数为准。
276- - key\_quant\_mode和value\_quant\_mode取值需要保持一致。277+ - key\_quant\_mode和value\_quant\_mode取值需要保持一致。
277- - dequant\_scale\_key和dequant\_scale\_value要么都为空,要么都不为空;dequant\_offset\_key和dequant\_offset\_value要么都为空,要么都不为空。278+ - dequant\_scale\_key和dequant\_scale\_value要么都为空,要么都不为空;dequant\_offset\_key和dequant\_offset\_value要么都为空,要么都不为空。
278- - dequant\_scale\_key和dequant\_scale\_value都不为空时,其shape需要保持一致;dequant\_offset\_key和dequant\_offset\_value都不为空时,其shape需要保持一致。 279+ - dequant\_scale\_key和dequant\_scale\_value都不为空时,其shape需要保持一致;dequant\_offset\_key和dequant\_offset\_value都不为空时,其shape需要保持一致。
279- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>: 280+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
280- - 仅支持pertoken和perchannel模式,pertoken模式下要求两个参数的shape均为\(B, KV\_S\),数据类型固定为float32;perchannel模式下要求两个参数的shape为(KV\_N, D),\(KV\_N, D\),\(H\),数据类型固定为bfloat16,H为KV\_N*D。281+ - 仅支持pertoken和perchannel模式,pertoken模式下要求两个参数的shape均为\(B, KV\_S\),数据类型固定为float32;perchannel模式下要求两个参数的shape为(KV\_N, D),\(KV\_N, D\),\(H\),数据类型固定为bfloat16,H为KV\_N*D。
281- - dequant\_scale\_key与dequant\_scale\_value非空场景,要求query的s小于等于16;要求query的dtype为bfloat16,key、value的dtype为int8,输出的dtype为bfloat16;不支持tensorlist、page attention特性。282+ - dequant\_scale\_key与dequant\_scale\_value非空场景,要求query的s小于等于16;要求query的dtype为bfloat16,key、value的dtype为int8,输出的dtype为bfloat16;不支持tensorlist、page attention特性。
282 283
283- - 管理scale/offset的量化模式如下:284+ - 管理scale/offset的量化模式如下:
284 285
285 > [!NOTE] 286 > [!NOTE]
286 > 注意scale、offset具体指dequant\_scale\_key、dequant\_scale\_key、dequant\_offset\_value、dequant\_offset\_value参数。287 > 注意scale、offset具体指dequant\_scale\_key、dequant\_scale\_key、dequant\_offset\_value、dequant\_offset\_value参数。
@@ -315,41 +316,41 @@ torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=No
315 </tbody>316 </tbody>
316 </table>317 </table>
317 318
318-- **当Q\_S等于1时:**319+- **当Q\_S等于1时:**
319- - query、key、value输入,功能使用限制如下:320+ - query、key、value输入,功能使用限制如下:
320- - 支持B轴小于等于65536,支持N轴小于等于256,支持S轴小于等于262144,支持D轴小于等于512。321+ - 支持B轴小于等于65536,支持N轴小于等于256,支持S轴小于等于262144,支持D轴小于等于512。
321- - query、key、value输入类型均为int8的场景暂不支持。322+ - query、key、value输入类型均为int8的场景暂不支持。
322- - 在int4(int32)伪量化场景下,PyTorch入图调用仅支持KV int4拼接成int32输入(建议通过dynamicQuant生成int4格式的数据,因为dynamicQuant就是一个int32包括8个int4)。323+ - 在int4(int32)伪量化场景下,PyTorch入图调用仅支持KV int4拼接成int32输入(建议通过dynamicQuant生成int4格式的数据,因为dynamicQuant就是一个int32包括8个int4)。
323- - 在int4(int32)伪量化场景下,若KV int4拼接成int32输入,那么KV的N、D或者H是实际值的八分之一。并且,int4伪量化仅支持D 64对齐(int32支持D 8对齐)。324+ - 在int4(int32)伪量化场景下,若KV int4拼接成int32输入,那么KV的N、D或者H是实际值的八分之一。并且,int4伪量化仅支持D 64对齐(int32支持D 8对齐)。
324 325 
325- - actual\_seq\_qlen:326+ - actual\_seq\_qlen:
326 327
327- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当query的input\_layout不为TND时,Q\_S为1时该参数无效。当query的input\_layout为TND/TND\_NTD时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。328+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当query的input\_layout不为TND时,Q\_S为1时该参数无效。当query的input\_layout为TND/TND\_NTD时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
328 329
329- - actual\_seq\_kvlen:330+ - actual\_seq\_kvlen:
330 331
331- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于key/value中对应batch的Sequence Length。seqlenKv的传入长度为1时,每个Batch使用相同seqlenKv;传入长度大于等于Batch时取seqlenKv的前Batch个数。其他长度不支持。当key/value的input\_layout为TND/TND\_NTD时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。332+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该入参中每个batch的有效Sequence Length应该不大于key/value中对应batch的Sequence Length。seqlenKv的传入长度为1时,每个Batch使用相同seqlenKv;传入长度大于等于Batch时取seqlenKv的前Batch个数。其他长度不支持。当key/value的input\_layout为TND/TND\_NTD时,综合约束请见[约束说明](#zh-cn_topic_0000001832267082_section12345537164214)。
332 333
333- - page attention场景:334+ - page attention场景:
334- - 使能必要条件是block\_table存在且有效,同时key、value是按照block\_table中的索引在一片连续内存中排布,在该场景下key、value的input\_layout参数无效。335+ - 使能必要条件是block\_table存在且有效,同时key、value是按照block\_table中的索引在一片连续内存中排布,在该场景下key、value的input\_layout参数无效。
335- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:336+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
336- - 支持key、value数据类型为float16、bfloat16、int8。337+ - 支持key、value数据类型为float16、bfloat16、int8。
337- - 不支持Q为`bfloat16`、float16、key、value为int4(int32)的场景。338+ - 不支持Q为`bfloat16`、float16、key、value为int4(int32)的场景。
338 339
339- - 该场景下,block\_size是用户自定义的参数,该参数的取值会影响page attention的性能。key、value输入类型为float16、bfloat16时需要16对齐,key、value输入类型为int8时需要32对齐,推荐使用128。通常情况下,page attention可以提高吞吐量,但会带来性能上的下降。340+ - 该场景下,block\_size是用户自定义的参数,该参数的取值会影响page attention的性能。key、value输入类型为float16、bfloat16时需要16对齐,key、value输入类型为int8时需要32对齐,推荐使用128。通常情况下,page attention可以提高吞吐量,但会带来性能上的下降。
340- - 参数key、value各自对应tensor的shape所有维度相乘不能超过int32的表示范围。341+ - 参数key、value各自对应tensor的shape所有维度相乘不能超过int32的表示范围。
341- - page attention场景下,blockTable必须为二维,第一维长度需等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为不同batch中最大actual\_seq\_kvlen对应的block数量)。342+ - page attention场景下,blockTable必须为二维,第一维长度需等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为不同batch中最大actual\_seq\_kvlen对应的block数量)。
342- - page attention场景下,当query的input\_layout为BNSD、TND时,kv cache排布支持(blocknum, blocksize, H)和(blocknum, KV\_N, blocksize, D)两种格式,当query的input\_layout为BSH、BSND时,kv cache排布只支持(blocknum, blocksize, H)一种格式。blocknum不能小于根据actual\_seq\_kvlen和blockSize计算的每个batch的block数量之和。且key和value的shape需保证一致。343+ - page attention场景下,当query的input\_layout为BNSD、TND时,kv cache排布支持(blocknum, blocksize, H)和(blocknum, KV\_N, blocksize, D)两种格式,当query的input\_layout为BSH、BSND时,kv cache排布只支持(blocknum, blocksize, H)一种格式。blocknum不能小于根据actual\_seq\_kvlen和blockSize计算的每个batch的block数量之和。且key和value的shape需保证一致。
343- - page attention场景下,kv cache排布为(blocknum, KV\_N, blocksize, D)时性能通常优于kv cache排布为(blocknum, blocksize, H)时的性能,建议优先选择(blocknum, KV\_N, blocksize, D)格式。344+ - page attention场景下,kv cache排布为(blocknum, KV\_N, blocksize, D)时性能通常优于kv cache排布为(blocknum, blocksize, H)时的性能,建议优先选择(blocknum, KV\_N, blocksize, D)格式。
344- - page attention使能场景下,当输入kv cache排布格式为(blocknum, blocksize, H),且numKvHeads \* headDim 超过64k时,受硬件指令约束,会被拦截报错。可通过使能GQA(减小 numKvHeads)或调整kv cache排布格式为(blocknum, numKvHeads, blocksize, D)解决。345+ - page attention使能场景下,当输入kv cache排布格式为(blocknum, blocksize, H),且numKvHeads \* headDim 超过64k时,受硬件指令约束,会被拦截报错。可通过使能GQA(减小 numKvHeads)或调整kv cache排布格式为(blocknum, numKvHeads, blocksize, D)解决。
345- - page attention场景的参数key、value各自对应tensor的shape所有维度相乘不能超过int32的表示范围。346+ - page attention场景的参数key、value各自对应tensor的shape所有维度相乘不能超过int32的表示范围。
346 347 
347- - kv伪量化参数分离:348+ - kv伪量化参数分离:
348- - 除了key\_quant\_mode为0并且value\_quant\_mode为1的场景外,key\_quant\_mode和value\_quant\_mode取值需要保持一致。 349+ - 除了key\_quant\_mode为0并且value\_quant\_mode为1的场景外,key\_quant\_mode和value\_quant\_mode取值需要保持一致。
349- - dequant\_scale\_key和dequant\_scale\_value要么都为空,要么都不为空;dequant\_offset\_key和dequant\_offset\_value要么都为空,要么都不为空。350+ - dequant\_scale\_key和dequant\_scale\_value要么都为空,要么都不为空;dequant\_offset\_key和dequant\_offset\_value要么都为空,要么都不为空。
350- - dequant\_scale\_key和dequant\_scale\_value都不为空时,除了key\_quant\_mode为0并且value\_quant\_mode为1的场景外,其shape需要保持一致;dequant\_offset\_key和dequant\_offset\_value都不为空时,除了key\_quant\_mode为0并且value\_quant\_mode为1的场景外,其shape需要保持一致。351+ - dequant\_scale\_key和dequant\_scale\_value都不为空时,除了key\_quant\_mode为0并且value\_quant\_mode为1的场景外,其shape需要保持一致;dequant\_offset\_key和dequant\_offset\_value都不为空时,除了key\_quant\_mode为0并且value\_quant\_mode为1的场景外,其shape需要保持一致。
351- - int4(int32)伪量化场景不支持后量化。352+ - int4(int32)伪量化场景不支持后量化。
352- - 管理scale/offset的量化模式如下:353+ - 管理scale/offset的量化模式如下:
353 354
354 > [!NOTE] 355 > [!NOTE]
355 > 注意scale、offset两个参数指dequant\_scale\_key、dequant\_scale\_key、dequant\_offset\_value、dequant\_offset\_value。356 > 注意scale、offset两个参数指dequant\_scale\_key、dequant\_scale\_key、dequant\_offset\_value、dequant\_offset\_value。
@@ -448,16 +449,16 @@ torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=No
448 <tr id="zh-cn_topic_0000001832267082_row194748261012"><td class="cellrowborder" valign="top" headers="mcps1.1.5.1.1 "><p id="zh-cn_topic_0000001832267082_p1154111491113"><a name="zh-cn_topic_0000001832267082_p1154111491113"></a><a name="zh-cn_topic_0000001832267082_p1154111491113"></a>对于value支持pertoken,两个参数的shape均为(1, B, KV_S)并且数据类型固定为float32。</p>449 <tr id="zh-cn_topic_0000001832267082_row194748261012"><td class="cellrowborder" valign="top" headers="mcps1.1.5.1.1 "><p id="zh-cn_topic_0000001832267082_p1154111491113"><a name="zh-cn_topic_0000001832267082_p1154111491113"></a><a name="zh-cn_topic_0000001832267082_p1154111491113"></a>对于value支持pertoken,两个参数的shape均为(1, B, KV_S)并且数据类型固定为float32。</p>
449 </td>450 </td>
450 </tr>451 </tr>
451- </tbody>452+ </tbody>
452 </table>453 </table>
453 454
454- - pse\_shift功能使用限制如下:455+ - pse\_shift功能使用限制如下:
455- - pse\_shift数据类型需与query数据类型保持一致。456+ - pse\_shift数据类型需与query数据类型保持一致。
456- - 仅支持D轴对齐,即D轴可以被16整除。457+ - 仅支持D轴对齐,即D轴可以被16整除。
457 458 
458## 调用示例<a name="zh-cn_topic_0000001832267082_section14459801435"></a>459## 调用示例<a name="zh-cn_topic_0000001832267082_section14459801435"></a>
459 460 
460-- 单算子模式调用461+- 单算子模式调用
461 462 
462 ```python463 ```python
463 import torch464 import torch
@@ -487,7 +488,7 @@ torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=No
487 device='npu:0', dtype=torch.float16)488 device='npu:0', dtype=torch.float16)
488 ```489 ```
489 490 
490-- 图模式调用491+- 图模式调用
491 492 
492 ```python493 ```python
493 import torch494 import torch
@@ -550,4 +551,3 @@ torch_npu.npu_fused_infer_attention_score_v2(query, key, value, *, query_rope=No
550 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],551 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],
551 device='npu:0', dtype=torch.float16) torch.Size([1, 8, 164, 128])552 device='npu:0', dtype=torch.float16) torch.Size([1, 8, 164, 128])
552 ```553 ```
553- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_fusion_attention.md+83-84
@@ -14,42 +14,42 @@
14 14 
15## 函数原型<a name="zh-cn_topic_0000001742717129_section45077510411"></a>15## 函数原型<a name="zh-cn_topic_0000001742717129_section45077510411"></a>
16 16 
17-```17+```python
18torch_npu.npu_fusion_attention(query, key, value, head_num, input_layout, pse=None, padding_mask=None, atten_mask=None, scale=1., keep_prob=1., pre_tockens=2147483647, next_tockens=2147483647, inner_precise=0, prefix=None, actual_seq_qlen=None, actual_seq_kvlen=None, sparse_mode=0, gen_mask_parallel=True, sync=False, softmax_layout="", sink=None, dropout_mask=None, seed=0, offset=0) -> (Tensor, Tensor, Tensor, Tensor, int, int, int)18torch_npu.npu_fusion_attention(query, key, value, head_num, input_layout, pse=None, padding_mask=None, atten_mask=None, scale=1., keep_prob=1., pre_tockens=2147483647, next_tockens=2147483647, inner_precise=0, prefix=None, actual_seq_qlen=None, actual_seq_kvlen=None, sparse_mode=0, gen_mask_parallel=True, sync=False, softmax_layout="", sink=None, dropout_mask=None, seed=0, offset=0) -> (Tensor, Tensor, Tensor, Tensor, int, int, int)
19```19```
20 20 
21## 参数说明<a name="zh-cn_topic_0000001742717129_section112637109429"></a>21## 参数说明<a name="zh-cn_topic_0000001742717129_section112637109429"></a>
22 22 
23-- **query**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。23+- **query**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
24-- **key**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。24+- **key**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
25-- **value**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。25+- **value**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
26-- **head\_num**(`int`):代表head个数,数据类型支持`int64`。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。26+- **head\_num**(`int`):代表head个数,数据类型支持`int64`。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
27-- **input\_layout**(`string`):代表输入`query`、`key`、`value`的数据排布格式,支持BSH、SBH、BSND、BNSD、TND(`actual_seq_qlen`/`actual_seq_kvlen`需传值,`input_layout`为TND时即为varlen场景);后续章节如无特殊说明,S表示`query`或`key`、`value`的sequence length,Sq表示query的sequence length,Skv表示`key`、`value`的sequence length,SS表示Sq\*Skv。27+- **input\_layout**(`string`):代表输入`query`、`key`、`value`的数据排布格式,支持BSH、SBH、BSND、BNSD、TND(`actual_seq_qlen`/`actual_seq_kvlen`需传值,`input_layout`为TND时即为varlen场景);后续章节如无特殊说明,S表示`query`或`key`、`value`的sequence length,Sq表示query的sequence length,Skv表示`key`、`value`的sequence length,SS表示Sq\*Skv。
28-- **pse**(`Tensor`):可选参数,表示位置编码。数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。28+- **pse**(`Tensor`):可选参数,表示位置编码。数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。
29 - 非varlen场景支持四维输入,包含BNSS格式、BN1Skv格式、1NSS格式。29 - 非varlen场景支持四维输入,包含BNSS格式、BN1Skv格式、1NSS格式。
30 - 若非varlen场景Sq大于1024或varlen场景、每个batch的Sq与Skv等长且是sparse\_mode为0、2、3的下三角掩码场景,可使能alibi位置编码压缩,此时只需要输入原始PSE最后1024行进行内存优化,即alibi\_compress = ori\_pse\[:, :, -1024:, :\],参数每个batch不相同时,输入BNHSkv\(H=1024\),每个batch相同时,输入1NHSkv\(H=1024\)。30 - 若非varlen场景Sq大于1024或varlen场景、每个batch的Sq与Skv等长且是sparse\_mode为0、2、3的下三角掩码场景,可使能alibi位置编码压缩,此时只需要输入原始PSE最后1024行进行内存优化,即alibi\_compress = ori\_pse\[:, :, -1024:, :\],参数每个batch不相同时,输入BNHSkv\(H=1024\),每个batch相同时,输入1NHSkv\(H=1024\)。
31-- **padding\_mask**(`Tensor`):暂不支持该传参。31+- **padding\_mask**(`Tensor`):暂不支持该传参。
32-- **atten\_mask**(`Tensor`):可选参数,取值为1代表该位不参与计算(不生效),为0代表该位参与计算,数据类型支持`bool`、`uint8`,数据格式支持$ND$,输入shape类型支持BNSS格式、B1SS格式、11SS格式、SS格式。varlen场景只支持SS格式,SS分别是maxSq和maxSkv。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。32+- **atten\_mask**(`Tensor`):可选参数,取值为1代表该位不参与计算(不生效),为0代表该位参与计算,数据类型支持`bool`、`uint8`,数据格式支持$ND$,输入shape类型支持BNSS格式、B1SS格式、11SS格式、SS格式。varlen场景只支持SS格式,SS分别是maxSq和maxSkv。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
33-- **scale**(`float`):可选参数,代表缩放系数,作为计算流中Muls的scalar值,数据类型支持`float`,默认值为1。33+- **scale**(`float`):可选参数,代表缩放系数,作为计算流中Muls的scalar值,数据类型支持`float`,默认值为1。
34-- **keep\_prob**(`float`):可选参数,代表Dropout中1的比例,取值范围为\(0, 1\] 。数据类型支持`float`,默认值为1,表示全部保留。34+- **keep\_prob**(`float`):可选参数,代表Dropout中1的比例,取值范围为\(0, 1\] 。数据类型支持`float`,默认值为1,表示全部保留。
35-- **pre\_tockens**(`int`):用于稀疏计算的参数,可选参数,数据类型支持`int64`,默认值为2147483647。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。35+- **pre\_tockens**(`int`):用于稀疏计算的参数,可选参数,数据类型支持`int64`,默认值为2147483647。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
36-- **next\_tockens**(`int`):用于稀疏计算的参数,可选参数,数据类型支持`int64`,默认值为2147483647。`next_tockens`和`pre_tockens`取值与`atten_mask`的关系请参见`sparse_mode`参数,参数取值与`atten_mask`分布不一致会导致精度问题。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。36+- **next\_tockens**(`int`):用于稀疏计算的参数,可选参数,数据类型支持`int64`,默认值为2147483647。`next_tockens`和`pre_tockens`取值与`atten_mask`的关系请参见`sparse_mode`参数,参数取值与`atten_mask`分布不一致会导致精度问题。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
37-- **inner\_precise**(`int`):用于提升精度,数据类型支持`int64`,默认值为0。37+- **inner\_precise**(`int`):用于提升精度,数据类型支持`int64`,默认值为0。
38 38 
39 > [!NOTE] 39 > [!NOTE]
40 > 当前0、1为保留配置值,2为使能无效行计算,其功能是避免在计算过程中存在整行mask进而导致精度有损失,但是该配置会导致性能下降。40 > 当前0、1为保留配置值,2为使能无效行计算,其功能是避免在计算过程中存在整行mask进而导致精度有损失,但是该配置会导致性能下降。
41 >如果算子可判断出存在无效行场景,会自动使能无效行计算,例如`sparse_mode`为3,Sq \> Skv场景。41 >如果算子可判断出存在无效行场景,会自动使能无效行计算,例如`sparse_mode`为3,Sq \> Skv场景。
42 42 
43-- **prefix**(`List[int]`):可选参数,代表prefix稀疏计算场景每个Batch的N值。数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。43+- **prefix**(`List[int]`):可选参数,代表prefix稀疏计算场景每个Batch的N值。数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
44-- **actual\_seq\_qlen**(`List[int]`):可选参数,varlen场景时需要传入此参数。表示`query`每个S的累加和长度,数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。44+- **actual\_seq\_qlen**(`List[int]`):可选参数,varlen场景时需要传入此参数。表示`query`每个S的累加和长度,数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
45 45 
46 比如真正的S长度列表为:2 2 2 2 2,则`actual_seq_qlen`传:2 4 6 8 10。46 比如真正的S长度列表为:2 2 2 2 2,则`actual_seq_qlen`传:2 4 6 8 10。
47 47 
48-- **actual\_seq\_kvlen**(`List[int]`):可选参数,varlen场景时需要传入此参数。表示`key`/`value`每个S的累加和长度。数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。48+- **actual\_seq\_kvlen**(`List[int]`):可选参数,varlen场景时需要传入此参数。表示`key`/`value`每个S的累加和长度。数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
49 49 
50 比如真正的S长度列表为:2 2 2 2 2,则actual\_seq\_kvlen传:2 4 6 8 10。50 比如真正的S长度列表为:2 2 2 2 2,则actual\_seq\_kvlen传:2 4 6 8 10。
51 51 
52-- **sparse\_mode**(`int`):表示sparse的模式,可选参数,默认值为0。取值如[表1](#zh-cn_topic_0000001742717129_table1946917414436)所示,不同模式的原理参见[参考资源](#zh-cn_topic_0000001742717129_section28169228374)。当整网的`atten_mask`都相同且shape小于2048\*2048时,建议使用defaultMask模式,来减少内存使用量。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。52+- **sparse\_mode**(`int`):表示sparse的模式,可选参数,默认值为0。取值如[表1](#zh-cn_topic_0000001742717129_table1946917414436)所示,不同模式的原理参见[参考资源](#zh-cn_topic_0000001742717129_section28169228374)。当整网的`atten_mask`都相同且shape小于2048\*2048时,建议使用defaultMask模式,来减少内存使用量。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
53 53 
54 **表1** sparse\_mode不同取值场景说明54 **表1** sparse\_mode不同取值场景说明
55 55 
@@ -128,57 +128,57 @@ torch_npu.npu_fusion_attention(query, key, value, head_num, input_layout, pse=No
128 </tbody>128 </tbody>
129 </table>129 </table>
130 130 
131-- **gen\_mask\_parallel**(`bool`):DSA生成dropout随机数向量mask的控制开关。默认值为True:同AI Core并行计算;设为False:同AI Core串行计算。131+- **gen\_mask\_parallel**(`bool`):DSA生成dropout随机数向量mask的控制开关。默认值为True:同AI Core并行计算;设为False:同AI Core串行计算。
132-- **sync**(`bool`):DSA生成dropout随机数向量mask的控制开关。默认值为False:dropout mask异步生成;设为True:dropout mask同步生成。132+- **sync**(`bool`):DSA生成dropout随机数向量mask的控制开关。默认值为False:dropout mask异步生成;设为True:dropout mask同步生成。
133-- **softmax_layout**(`string`):可选参数,用于控制TND场景下softmax的输出(softmax_max和softmax_sum)的数据排布方式。当前仅在input_layout=“TND”时进行配置,仅支持传入“TND”。默认情况下,softmax的输出排布为NTD排布;传入TND时,softmax的输出排布为TND排布。此参数为Ascend Extension for PyTorch 7.2.0版本新增参数,支持在CANN8.3.RC1及以上版本使用。133+- **softmax_layout**(`string`):可选参数,用于控制TND场景下softmax的输出(softmax_max和softmax_sum)的数据排布方式。当前仅在input_layout=“TND”时进行配置,仅支持传入“TND”。默认情况下,softmax的输出排布为NTD排布;传入TND时,softmax的输出排布为TND排布。此参数为Ascend Extension for PyTorch 7.2.0版本新增参数,支持在CANN8.3.RC1及以上版本使用。
134-- **sink**(`Tensor`):可选参数,每个注意力头的偏置。shape为`[head_num]`,数据类型仅支持`float32`。此参数为Ascend Extension for PyTorch 7.3.0版本新增参数,支持在CANN8.5.0及以上版本使用。134+- **sink**(`Tensor`):可选参数,每个注意力头的偏置。shape为`[head_num]`,数据类型仅支持`float32`。此参数为Ascend Extension for PyTorch 7.3.0版本新增参数,支持在CANN8.5.0及以上版本使用。
135-- **dropout\_mask**(`Tensor`):可选参数,外部传入的dropout掩码,用于控制dropout行为。当传入此参数时,将使用外部掩码而非内部生成,实现可复现的dropout效果。数据类型支持`uint8`,数据格式支持$ND$。若不传入此参数,则由算子内部根据`seed`和`offset`自动生成dropout mask。若传入此参数,则需要使用内部接口[_npu_dropout_gen_mask](https://gitcode.com/Ascend/op-plugin/blob/master/op_plugin/ops/opapi/DropoutGenMaskKernelNpuOpApi.cpp)生成dropout_mask。135+- **dropout\_mask**(`Tensor`):可选参数,外部传入的dropout掩码,用于控制dropout行为。当传入此参数时,将使用外部掩码而非内部生成,实现可复现的dropout效果。数据类型支持`uint8`,数据格式支持$ND$。若不传入此参数,则由算子内部根据`seed`和`offset`自动生成dropout mask。若传入此参数,则需要使用内部接口[_npu_dropout_gen_mask](https://gitcode.com/Ascend/op-plugin/blob/master/op_plugin/ops/opapi/DropoutGenMaskKernelNpuOpApi.cpp)生成dropout_mask。
136-- **seed**(`int`):可选参数,DSA生成dropout mask中Philox算法的种子值,数据类型支持`int64`,默认值为0。当`dropout_mask`参数未传入时,若`seed`为0,则使用默认随机种子;若`seed`非0,则使用指定的种子值生成dropout mask。当传入dropout_mask时,需传入生成dropout_mask所需的seed值。136+- **seed**(`int`):可选参数,DSA生成dropout mask中Philox算法的种子值,数据类型支持`int64`,默认值为0。当`dropout_mask`参数未传入时,若`seed`为0,则使用默认随机种子;若`seed`非0,则使用指定的种子值生成dropout mask。当传入dropout_mask时,需传入生成dropout_mask所需的seed值。
137-- **offset**(`int`):可选参数,DSA生成dropout mask中Philox算法的偏移值,数据类型支持`int64`,默认值为0。配合`seed`参数使用,用于控制dropout mask的生成位置。当传入dropout_mask时,需传入生成dropout_mask所需的offset值。137+- **offset**(`int`):可选参数,DSA生成dropout mask中Philox算法的偏移值,数据类型支持`int64`,默认值为0。配合`seed`参数使用,用于控制dropout mask的生成位置。当传入dropout_mask时,需传入生成dropout_mask所需的offset值。
138 138 
139## 输出说明<a name="zh-cn_topic_0000001742717129_section22231435517"></a>139## 输出说明<a name="zh-cn_topic_0000001742717129_section22231435517"></a>
140 140 
141共7个输出,类型依次为**Tensor、Tensor、Tensor、Tensor、int、int、int。**141共7个输出,类型依次为**Tensor、Tensor、Tensor、Tensor、int、int、int。**
142 142 
143-- 第1个输出为`Tensor`,计算公式的最终输出$attention\_out$,数据类型支持`float16`、`bfloat16`、`float32`。143+- 第1个输出为`Tensor`,计算公式的最终输出$attention\_out$,数据类型支持`float16`、`bfloat16`、`float32`。
144-- 第2个输出为`Tensor`,Softmax计算的Max中间结果,用于反向计算,数据类型支持`float`。144+- 第2个输出为`Tensor`,Softmax计算的Max中间结果,用于反向计算,数据类型支持`float`。
145-- 第3个输出为`Tensor`,Softmax计算的Sum中间结果,用于反向计算,数据类型支持`float`。145+- 第3个输出为`Tensor`,Softmax计算的Sum中间结果,用于反向计算,数据类型支持`float`。
146-- 第4个输出为`Tensor`,预留参数,暂未使用。146+- 第4个输出为`Tensor`,预留参数,暂未使用。
147-- 第5个输出为`int`,DSA生成dropout mask中,Philox算法的seed。147+- 第5个输出为`int`,DSA生成dropout mask中,Philox算法的seed。
148-- 第6个输出为`int`,DSA生成dropout mask中,Philox算法的offset。148+- 第6个输出为`int`,DSA生成dropout mask中,Philox算法的offset。
149-- 第7个输出为`int`,DSA生成dropout mask的长度。149+- 第7个输出为`int`,DSA生成dropout mask的长度。
150 150 
151## 约束说明<a name="zh-cn_topic_0000001742717129_section12345537164214"></a>151## 约束说明<a name="zh-cn_topic_0000001742717129_section12345537164214"></a>
152 152 
153-- 该接口仅在训练场景下使用。153+- 该接口仅在训练场景下使用。
154-- 该接口暂不支持图模式,不支持aclgraph。154+- 该接口暂不支持图模式,不支持aclgraph。
155-- 输入`query`、`key`、`value`、`pse`的数据类型必须一致。155+- 输入`query`、`key`、`value`、`pse`的数据类型必须一致。
156-- 输入`query`、`key`、`value`的`input_layout`必须一致。156+- 输入`query`、`key`、`value`的`input_layout`必须一致。
157-- 输入`query`、`key`、`value`的shape说明:157+- 输入`query`、`key`、`value`的shape说明:
158- - 输入`key`和`value`的shape必须一致。158+ - 输入`key`和`value`的shape必须一致。
159- - B:batchsize必须相等;非varlen场景B取值范围1\~2M;varlen场景B取值范围1\~2K。159+ - B:batchsize必须相等;非varlen场景B取值范围1\~2M;varlen场景B取值范围1\~2K。
160- - D:Head Dim必须满足Dq=Dk和Dk≥Dv,取值范围1\~768。160+ - D:Head Dim必须满足Dq=Dk和Dk≥Dv,取值范围1\~768。
161- - S:sequence length,取值范围1\~1M。161+ - S:sequence length,取值范围1\~1M。
162 162 
163-- varlen场景下:163+- varlen场景下:
164- - 要求T(B\*S)取值范围1\~1M。164+ - 要求T(B\*S)取值范围1\~1M。
165- - `atten_mask`输入不支持补pad,即`atten_mask`中不能存在某一行全1的场景。165+ - `atten_mask`输入不支持补pad,即`atten_mask`中不能存在某一行全1的场景。
166 166 
167-- 支持输入`query`的N和`key`/`value`的N不相等,但必须成比例关系,即Nq/Nkv必须是非0整数,Nq取值范围1\~256。当Nq/Nkv \> 1时,即为GQA\(grouped-query attention\);当Nq/Nkv=1时,即为MHA\(multi-head attention\)。167+- 支持输入`query`的N和`key`/`value`的N不相等,但必须成比例关系,即Nq/Nkv必须是非0整数,Nq取值范围1\~256。当Nq/Nkv \> 1时,即为GQA\(grouped-query attention\);当Nq/Nkv=1时,即为MHA\(multi-head attention\)。
168 168 
169 > [!NOTE] 169 > [!NOTE]
170 > 本文如无特殊说明,N表示的是Nq。170 > 本文如无特殊说明,N表示的是Nq。
171 171 
172-- `sparse_mode`取值说明:172+- `sparse_mode`取值说明:
173- - `sparse_mode`为1、2、3、4、5、6、7、8时,应传入对应正确的`atten_mask`,否则将导致计算结果错误。当`atten_mask`输入为None时,`sparse_mode`,`pre_tockens`,`next_tockens`参数不生效,固定为全计算。173+ - `sparse_mode`为1、2、3、4、5、6、7、8时,应传入对应正确的`atten_mask`,否则将导致计算结果错误。当`atten_mask`输入为None时,`sparse_mode`,`pre_tockens`,`next_tockens`参数不生效,固定为全计算。
174- - `sparse_mode`配置为1、2、3、5、6时,用户配置的`pre_tockens`、`next_tockens`不会生效。174+ - `sparse_mode`配置为1、2、3、5、6时,用户配置的`pre_tockens`、`next_tockens`不会生效。
175- - `sparse_mode`配置为0、4时,需保证`atten_mask`与`pre_tockens`、`next_tockens`的范围一致。175+ - `sparse_mode`配置为0、4时,需保证`atten_mask`与`pre_tockens`、`next_tockens`的范围一致。
176- - `sparse_mode`配置为7或者8时,不支持可选参数`pse`。176+ - `sparse_mode`配置为7或者8时,不支持可选参数`pse`。
177 177 
178-- `prefix`稀疏计算场景B不大于32,varlen场景不支持非压缩prefix,即不支持sparse\_mode=5;当Sq\>Skv时,`prefix`的N值取值范围\[0, Skv\],当Sq<=Skv时,`prefix`的N值取值范围\[Skv-Sq, Skv\]。178+- `prefix`稀疏计算场景B不大于32,varlen场景不支持非压缩prefix,即不支持sparse\_mode=5;当Sq\>Skv时,`prefix`的N值取值范围\[0, Skv\],当Sq<=Skv时,`prefix`的N值取值范围\[Skv-Sq, Skv\]。
179-- 支持`actual_seq_qlen`中某个Batch上的S长度为0;如果存在S为0的情况,不支持`pse`输入,假设真实的S长度为\[2, 2, 0, 2, 2\],则传入的`actual_seq_qlen`为\[2, 4, 4, 6, 8\]。`actual_seq_qlen`的长度取值范围为1\~2K,varlen场景下长度最大支持1K。179+- 支持`actual_seq_qlen`中某个Batch上的S长度为0;如果存在S为0的情况,不支持`pse`输入,假设真实的S长度为\[2, 2, 0, 2, 2\],则传入的`actual_seq_qlen`为\[2, 4, 4, 6, 8\]。`actual_seq_qlen`的长度取值范围为1\~2K,varlen场景下长度最大支持1K。
180-- TND格式下,支持尾部部分Batch不参与计算,此时`actual_seq_qlen`和`actual_seq_kv_len`尾部传入对应个数个0即可。假设真实的S长度为\[2, 3, 4, 5, 6\],此时后两个Batch不参与计算,则传入的`actual_seq_qlen`为\[2, 5, 9, 0, 0\]。180+- TND格式下,支持尾部部分Batch不参与计算,此时`actual_seq_qlen`和`actual_seq_kv_len`尾部传入对应个数个0即可。假设真实的S长度为\[2, 3, 4, 5, 6\],此时后两个Batch不参与计算,则传入的`actual_seq_qlen`为\[2, 5, 9, 0, 0\]。
181-- 部分场景下,如果计算量过大可能会导致算子执行超时\(aicore error类型报错,errorStr为:timeout or trap error\),此时建议做轴切分处理,注:这里的计算量会受B、S、N、D等参数的影响,值越大计算量越大。181+- 部分场景下,如果计算量过大可能会导致算子执行超时\(aicore error类型报错,errorStr为:timeout or trap error\),此时建议做轴切分处理,注:这里的计算量会受B、S、N、D等参数的影响,值越大计算量越大。
182 182 
183## 调用示例<a name="zh-cn_topic_0000001742717129_section14459801435"></a>183## 调用示例<a name="zh-cn_topic_0000001742717129_section14459801435"></a>
184 184 
@@ -277,7 +277,7 @@ if __name__ == "__main__":
277```277```
278 278 
279使用外部dropout\_mask的示例:279使用外部dropout\_mask的示例:
280- 280+ 
281```python281```python
282import torch282import torch
283import torch_npu283import torch_npu
@@ -340,12 +340,12 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
340 340 
341**说明:下图中的蓝色表示保留该值,`atten_mask`中,应该配置为False;阴影表示遮蔽该值,`atten_mask`中应配置为True。**341**说明:下图中的蓝色表示保留该值,`atten_mask`中,应该配置为False;阴影表示遮蔽该值,`atten_mask`中应配置为True。**
342 342 
343-- 当`sparse_mode`为0时,代表defaultMask模式。343+- 当`sparse_mode`为0时,代表defaultMask模式。
344- - 不传mask:如果`atten_mask`未传入则不做mask操作,`atten_mask`取值为None,忽略`pre_tockens`和`next_tockens`取值。Masked QK<sup>T</sup>矩阵示意如下:344+ - 不传mask:如果`atten_mask`未传入则不做mask操作,`atten_mask`取值为None,忽略`pre_tockens`和`next_tockens`取值。Masked QK<sup>T</sup>矩阵示意如下:
345 345 
346 ![](../../figures/mode0.png)346 ![](../../figures/mode0.png)
347 347 
348- - `next_tockens`取值为0,`pre_tockens`大于等于Sq,表示causal场景sparse,`atten_mask`应传入下三角矩阵,此时`pre_tockens`和`next_tockens`之间的部分需要计算,Masked QK<sup>T</sup>矩阵示意如下:348+ - `next_tockens`取值为0,`pre_tockens`大于等于Sq,表示causal场景sparse,`atten_mask`应传入下三角矩阵,此时`pre_tockens`和`next_tockens`之间的部分需要计算,Masked QK<sup>T</sup>矩阵示意如下:
349 349 
350 ![](../../figures/1.png)350 ![](../../figures/1.png)
351 351 
@@ -353,7 +353,7 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
353 353 
354 ![](../../figures/1-1.png)354 ![](../../figures/1-1.png)
355 355 
356- - `pre_tockens`小于Sq,`next_tockens`小于Skv,且都大于等于0,表示band场景,此时`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>矩阵示意如下:356+ - `pre_tockens`小于Sq,`next_tockens`小于Skv,且都大于等于0,表示band场景,此时`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>矩阵示意如下:
357 357 
358 ![](../../figures/1-2.png)358 ![](../../figures/1-2.png)
359 359 
@@ -361,25 +361,25 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
361 361 
362 ![](../../figures/1-3.png)362 ![](../../figures/1-3.png)
363 363 
364- - `next_tockens`为负数,以pre\_tockens=9,next\_tockens=-3为例,`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>示意如下:364+ - `next_tockens`为负数,以pre\_tockens=9,next\_tockens=-3为例,`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>示意如下:
365 365 
366 **说明:`next_tockens`为负数时,`pre_tockens`取值必须大于等于`next_tockens`的绝对值,且`next_tockens`的绝对值小于Skv。**366 **说明:`next_tockens`为负数时,`pre_tockens`取值必须大于等于`next_tockens`的绝对值,且`next_tockens`的绝对值小于Skv。**
367 367 
368 ![](../../figures/1-4.png)368 ![](../../figures/1-4.png)
369 369 
370- - `pre_tockens`为负数,以next\_tockens=7,pre\_tockens=-3为例,`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>示意如下:370+ - `pre_tockens`为负数,以next\_tockens=7,pre\_tockens=-3为例,`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>示意如下:
371 371 
372 **说明:`pre_tockens`为负数时,`next_tockens`取值必须大于等于`pre_tockens`的绝对值,且`pre_tockens`的绝对值小于Sq。**372 **说明:`pre_tockens`为负数时,`next_tockens`取值必须大于等于`pre_tockens`的绝对值,且`pre_tockens`的绝对值小于Sq。**
373 373 
374 ![](../../figures/1-5.png)374 ![](../../figures/1-5.png)
375 375 
376-- 当`sparse_mode`为1时,代表allMask,即传入完整的`atten_mask`矩阵。376+- 当`sparse_mode`为1时,代表allMask,即传入完整的`atten_mask`矩阵。
377 377 
378 该场景下忽略`next_tockens`、`pre_tockens`取值,Masked QK<sup>T</sup>矩阵示意如下:378 该场景下忽略`next_tockens`、`pre_tockens`取值,Masked QK<sup>T</sup>矩阵示意如下:
379 379 
380 ![](../../figures/1-6.png)380 ![](../../figures/1-6.png)
381 381 
382-- 当`sparse_mode`为2时,代表leftUpCausal模式的mask,对应以左上顶点划分的下三角场景(参数起点为左上角)。该场景下忽略`pre_tockens`、`next_tockens`取值,Masked QK<sup>T</sup>矩阵示意如下:382+- 当`sparse_mode`为2时,代表leftUpCausal模式的mask,对应以左上顶点划分的下三角场景(参数起点为左上角)。该场景下忽略`pre_tockens`、`next_tockens`取值,Masked QK<sup>T</sup>矩阵示意如下:
383 383 
384 ![](../../figures/1-7.png)384 ![](../../figures/1-7.png)
385 385 
@@ -387,15 +387,15 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
387 387 
388 ![](../../figures/1-8.png)388 ![](../../figures/1-8.png)
389 389 
390-- 当`sparse_mode`为3时,代表rightDownCausal模式的mask,对应以右下顶点划分的下三角场景(参数起点为右下角)。该场景下忽略`pre_tockens`、`next_tockens`取值。`atten_mask`为优化后的压缩下三角矩阵(2048\*2048),Masked QK<sup>T</sup>矩阵示意如下:390+- 当`sparse_mode`为3时,代表rightDownCausal模式的mask,对应以右下顶点划分的下三角场景(参数起点为右下角)。该场景下忽略`pre_tockens`、`next_tockens`取值。`atten_mask`为优化后的压缩下三角矩阵(2048\*2048),Masked QK<sup>T</sup>矩阵示意如下:
391 391 
392 ![](../../figures/1-9.png)392 ![](../../figures/1-9.png)
393 393 
394-- 当`sparse_mode`为4时,代表band场景,即计算`pre_tockens`和`next_tockens`之间的部分,参数起点为右下角,`pre_tockens`和`next_tockens`之间需要有交集。`atten_mask`为优化后的压缩下三角矩阵(2048\*2048)。Masked QK<sup>T</sup>矩阵示意如下:394+- 当`sparse_mode`为4时,代表band场景,即计算`pre_tockens`和`next_tockens`之间的部分,参数起点为右下角,`pre_tockens`和`next_tockens`之间需要有交集。`atten_mask`为优化后的压缩下三角矩阵(2048\*2048)。Masked QK<sup>T</sup>矩阵示意如下:
395 395 
396 ![](../../figures/1-10.png)396 ![](../../figures/1-10.png)
397 397 
398-- 当`sparse_mode`为5时,代表prefix非压缩场景,即在rightDownCausal的基础上,左侧加上一个长为Sq,宽为N的矩阵,N的值由可选参数prefix获取,例如下图中表示batch=2场景下prefix传入数组\[4,5\],每个batch轴的N值可以不一样,参数起点为左上角。398+- 当`sparse_mode`为5时,代表prefix非压缩场景,即在rightDownCausal的基础上,左侧加上一个长为Sq,宽为N的矩阵,N的值由可选参数prefix获取,例如下图中表示batch=2场景下prefix传入数组\[4,5\],每个batch轴的N值可以不一样,参数起点为左上角。
399 399 
400 该场景下忽略`pre_tockens`、`next_tockens`取值,`atten_mask`矩阵数据格式须为BNSS或B1SS,Masked QK<sup>T</sup>矩阵示意如下:400 该场景下忽略`pre_tockens`、`next_tockens`取值,`atten_mask`矩阵数据格式须为BNSS或B1SS,Masked QK<sup>T</sup>矩阵示意如下:
401 401 
@@ -405,39 +405,38 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
405 405 
406 ![](../../figures/1-12.png)406 ![](../../figures/1-12.png)
407 407 
408-- 当`sparse_mode`为6时,代表prefix压缩场景,即prefix场景时,attenMask为优化后的压缩下三角+矩形的矩阵(3072\*2048):其中上半部分\[2048,2048\]的下三角矩阵,下半部分为\[1024,2048\]的矩形矩阵,矩形矩阵左半部分全0,右半部分全1,`atten_mask`应传入矩阵示意如下。该场景下忽略`pre_tockens`、`next_tockens`取值。408+- 当`sparse_mode`为6时,代表prefix压缩场景,即prefix场景时,attenMask为优化后的压缩下三角+矩形的矩阵(3072\*2048):其中上半部分\[2048,2048\]的下三角矩阵,下半部分为\[1024,2048\]的矩形矩阵,矩形矩阵左半部分全0,右半部分全1,`atten_mask`应传入矩阵示意如下。该场景下忽略`pre_tockens`、`next_tockens`取值。
409 409 
410 ![](../../figures/1-13.png)410 ![](../../figures/1-13.png)
411 411 
412-- 当`sparse_mode`为7时,表示varlen且为长序列外切场景(即长序列在模型脚本中进行多卡切query的sequence length);用户需要确保外切前为使用sparse\_mode=3的场景;当前mode下用户需要设置`pre_tockens`和`next_tockens`(起点为右下顶点),且需要保证参数正确,否则会存在精度问题。412+- 当`sparse_mode`为7时,表示varlen且为长序列外切场景(即长序列在模型脚本中进行多卡切query的sequence length);用户需要确保外切前为使用sparse\_mode=3的场景;当前mode下用户需要设置`pre_tockens`和`next_tockens`(起点为右下顶点),且需要保证参数正确,否则会存在精度问题。
413 413 
414 Masked QK<sup>T</sup>矩阵示意如下,在第二个batch对`query`进行切分,`key`和`value`不切分,4x6的mask矩阵被切分成2x6和2x6的mask,分别在卡1和卡2上计算:414 Masked QK<sup>T</sup>矩阵示意如下,在第二个batch对`query`进行切分,`key`和`value`不切分,4x6的mask矩阵被切分成2x6和2x6的mask,分别在卡1和卡2上计算:
415 415 
416- - 卡1的最后一块mask为band类型的mask,配置pre\_tockens=6(保证大于等于最后一个Skv),next\_tockens=-2,`actual_seq_qlen`应传入\{3,5\},`actual_seq_kvlen`应传入\{3,9\}。416+ - 卡1的最后一块mask为band类型的mask,配置pre\_tockens=6(保证大于等于最后一个Skv),next\_tockens=-2,`actual_seq_qlen`应传入\{3,5\},`actual_seq_kvlen`应传入\{3,9\}。
417- - 卡2的mask类型切分后不变,`sparse_mode`为3,`actual_seq_qlen`应传入\{2,7,11\},`actual_seq_kvlen`应传入\{6,11,15\}。417+ - 卡2的mask类型切分后不变,`sparse_mode`为3,`actual_seq_qlen`应传入\{2,7,11\},`actual_seq_kvlen`应传入\{6,11,15\}。
418 418 
419 ![](../../figures/1-14.png)419 ![](../../figures/1-14.png)
420 420 
421 > [!NOTE] 421 > [!NOTE]
422- > - 如果配置sparse\_mode=7,但实际只存在一个batch,用户需按照band模式的要求来配置参数;sparse\_mode=7时,用户需要输入2048x2048的下三角mask作为该融合算子的输入。422+ > - 如果配置sparse\_mode=7,但实际只存在一个batch,用户需按照band模式的要求来配置参数;sparse\_mode=7时,用户需要输入2048x2048的下三角mask作为该融合算子的输入。
423- > - 基于sparse\_mode=3进行外切产生的band模式的sparse的参数应符合以下条件:423+ > - 基于sparse\_mode=3进行外切产生的band模式的sparse的参数应符合以下条件:
424- > - pre\_tockens \>= last\_Skv。424+ > - pre\_tockens \>= last\_Skv。
425- > - next\_tockens <= 0。425+ > - next\_tockens <= 0。
426- > - 当前模式下不支持可选输入pse。426+ > - 当前模式下不支持可选输入pse。
427 427 
428-- 当`sparse_mode`为8时,表示varlen且为长序列外切场景;用户需要确保外切前为使用sparse\_mode=2的场景;当前mode下用户需要设置`pre_tockens`和`next_tockens`(起点为右下顶点),且需要保证参数正确,否则会存在精度问题。428+- 当`sparse_mode`为8时,表示varlen且为长序列外切场景;用户需要确保外切前为使用sparse\_mode=2的场景;当前mode下用户需要设置`pre_tockens`和`next_tockens`(起点为右下顶点),且需要保证参数正确,否则会存在精度问题。
429 429 
430 Masked QK<sup>T</sup>矩阵示意如下,在第二个batch对`query`进行切分,`key`和`value`不切分,5x4的mask矩阵被切分成2x4和3x4的mask,分别在卡1和卡2上计算:430 Masked QK<sup>T</sup>矩阵示意如下,在第二个batch对`query`进行切分,`key`和`value`不切分,5x4的mask矩阵被切分成2x4和3x4的mask,分别在卡1和卡2上计算:
431 431 
432- - 卡1的mask类型切分后不变,`sparse_mode`为2,`actual_seq_qlen`应传入\{3,5\},`actual_seq_kvlen`应传入\{3,7\}。432+ - 卡1的mask类型切分后不变,`sparse_mode`为2,`actual_seq_qlen`应传入\{3,5\},`actual_seq_kvlen`应传入\{3,7\}。
433- - 卡2的第一块mask为band类型的mask,配置pre\_tockens=4(保证大于等于第一个Skv),next\_tockens=1,`actual_seq_qlen`应传入\{3,8,12\},`actual_seq_kvlen`应传入\{4,9,13\}。433+ - 卡2的第一块mask为band类型的mask,配置pre\_tockens=4(保证大于等于第一个Skv),next\_tockens=1,`actual_seq_qlen`应传入\{3,8,12\},`actual_seq_kvlen`应传入\{4,9,13\}。
434 434 
435 ![](../../figures/1-15.png)435 ![](../../figures/1-15.png)
436 436 
437 > [!NOTE] 437 > [!NOTE]
438- > - 如果配置sparse\_mode=8,但实际只存在一个batch,用户需按照band模式的要求来配置参数;sparse\_mode=8时,用户需要输入2048x2048的下三角mask作为该融合算子的输入。438+ > - 如果配置sparse\_mode=8,但实际只存在一个batch,用户需按照band模式的要求来配置参数;sparse\_mode=8时,用户需要输入2048x2048的下三角mask作为该融合算子的输入。
439- > - 基于sparse\_mode=2进行外切产生的band模式的sparse的参数应符合以下条件:439+ > - 基于sparse\_mode=2进行外切产生的band模式的sparse的参数应符合以下条件:
440- > - pre\_tockens \>= first\_Skv。440+ > - pre\_tockens \>= first\_Skv。
441- > - next\_tockens范围无约束,根据实际情况进行配置。441+ > - next\_tockens范围无约束,根据实际情况进行配置。
442- > - 当前模式下不支持可选输入pse。442+ > - 当前模式下不支持可选输入pse。
443- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_fusion_attention_v3.md+78-78
@@ -16,42 +16,42 @@
16 16 
17## 函数原型<a name="zh-cn_topic_0000001742717129_section45077510411"></a>17## 函数原型<a name="zh-cn_topic_0000001742717129_section45077510411"></a>
18 18 
19-```19+```python
20torch_npu.npu_fusion_attention_v3(query, key, value, head_num, input_layout, pse=None, padding_mask=None, atten_mask=None, scale=1., keep_prob=1., pre_tockens=2147483647, next_tockens=2147483647, inner_precise=0, prefix=None, actual_seq_qlen=None, actual_seq_kvlen=None, sparse_mode=0, gen_mask_parallel=True, sync=False, softmax_layout="", sink=None) -> (Tensor, Tensor, Tensor, Tensor, Tensor, Tensor)20torch_npu.npu_fusion_attention_v3(query, key, value, head_num, input_layout, pse=None, padding_mask=None, atten_mask=None, scale=1., keep_prob=1., pre_tockens=2147483647, next_tockens=2147483647, inner_precise=0, prefix=None, actual_seq_qlen=None, actual_seq_kvlen=None, sparse_mode=0, gen_mask_parallel=True, sync=False, softmax_layout="", sink=None) -> (Tensor, Tensor, Tensor, Tensor, Tensor, Tensor)
21```21```
22 22 
23## 参数说明<a name="zh-cn_topic_0000001742717129_section112637109429"></a>23## 参数说明<a name="zh-cn_topic_0000001742717129_section112637109429"></a>
24 24 
25-- **query**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。25+- **query**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
26-- **key**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。26+- **key**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
27-- **value**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。27+- **value**(`Tensor`):数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
28-- **head\_num**(`int`):代表head个数,数据类型支持`int64`。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。28+- **head\_num**(`int`):代表head个数,数据类型支持`int64`。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
29-- **input\_layout**(`string`):代表输入`query`、`key`、`value`的数据排布格式,支持BSH、SBH、BSND、BNSD、TND(`actual_seq_qlen`/`actual_seq_kvlen`需传值,`input_layout`为TND时即为varlen场景);后续章节如无特殊说明,S表示`query`或`key`、`value`的sequence length,Sq表示query的sequence length,Skv表示`key`、`value`的sequence length,SS表示Sq\*Skv。29+- **input\_layout**(`string`):代表输入`query`、`key`、`value`的数据排布格式,支持BSH、SBH、BSND、BNSD、TND(`actual_seq_qlen`/`actual_seq_kvlen`需传值,`input_layout`为TND时即为varlen场景);后续章节如无特殊说明,S表示`query`或`key`、`value`的sequence length,Sq表示query的sequence length,Skv表示`key`、`value`的sequence length,SS表示Sq\*Skv。
30-- **pse**(`Tensor`):可选参数,表示位置编码。数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。30+- **pse**(`Tensor`):可选参数,表示位置编码。数据类型支持`float16`、`bfloat16`、`float32`,数据格式支持$ND$。
31 - 非varlen场景支持四维输入,包含BNSS格式、BN1Skv格式、1NSS格式。31 - 非varlen场景支持四维输入,包含BNSS格式、BN1Skv格式、1NSS格式。
32 - 若非varlen场景Sq大于1024或varlen场景、每个batch的Sq与Skv等长且是sparse\_mode为0、2、3的下三角掩码场景,可使能alibi位置编码压缩,此时只需要输入原始PSE最后1024行进行内存优化,即alibi\_compress = ori\_pse\[:, :, -1024:, :\],参数每个batch不相同时,输入BNHSkv\(H=1024\),每个batch相同时,输入1NHSkv\(H=1024\)。32 - 若非varlen场景Sq大于1024或varlen场景、每个batch的Sq与Skv等长且是sparse\_mode为0、2、3的下三角掩码场景,可使能alibi位置编码压缩,此时只需要输入原始PSE最后1024行进行内存优化,即alibi\_compress = ori\_pse\[:, :, -1024:, :\],参数每个batch不相同时,输入BNHSkv\(H=1024\),每个batch相同时,输入1NHSkv\(H=1024\)。
33-- **padding\_mask**(`Tensor`):暂不支持该传参。33+- **padding\_mask**(`Tensor`):暂不支持该传参。
34-- **atten\_mask**(`Tensor`):可选参数,取值为1代表该位不参与计算(不生效),为0代表该位参与计算,数据类型支持`bool`、`uint8`,数据格式支持$ND$,输入shape类型支持BNSS格式、B1SS格式、11SS格式、SS格式。varlen场景只支持SS格式,SS分别是maxSq和maxSkv。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。34+- **atten\_mask**(`Tensor`):可选参数,取值为1代表该位不参与计算(不生效),为0代表该位参与计算,数据类型支持`bool`、`uint8`,数据格式支持$ND$,输入shape类型支持BNSS格式、B1SS格式、11SS格式、SS格式。varlen场景只支持SS格式,SS分别是maxSq和maxSkv。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
35-- **scale**(`float`):可选参数,代表缩放系数,作为计算流中Muls的scalar值,数据类型支持`float`,默认值为1。35+- **scale**(`float`):可选参数,代表缩放系数,作为计算流中Muls的scalar值,数据类型支持`float`,默认值为1。
36-- **keep\_prob**(`float`):可选参数,代表Dropout中1的比例,取值范围为\(0, 1\] 。数据类型支持`float`,默认值为1,表示全部保留。36+- **keep\_prob**(`float`):可选参数,代表Dropout中1的比例,取值范围为\(0, 1\] 。数据类型支持`float`,默认值为1,表示全部保留。
37-- **pre\_tockens**(`int`):用于稀疏计算的参数,可选参数,数据类型支持`int64`,默认值为2147483647。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。37+- **pre\_tockens**(`int`):用于稀疏计算的参数,可选参数,数据类型支持`int64`,默认值为2147483647。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
38-- **next\_tockens**(`int`):用于稀疏计算的参数,可选参数,数据类型支持`int64`,默认值为2147483647。`next_tockens`和`pre_tockens`取值与`atten_mask`的关系请参见`sparse_mode`参数,参数取值与`atten_mask`分布不一致会导致精度问题。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。38+- **next\_tockens**(`int`):用于稀疏计算的参数,可选参数,数据类型支持`int64`,默认值为2147483647。`next_tockens`和`pre_tockens`取值与`atten_mask`的关系请参见`sparse_mode`参数,参数取值与`atten_mask`分布不一致会导致精度问题。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
39-- **inner\_precise**(`int`):用于提升精度,数据类型支持`int64`,默认值为0。39+- **inner\_precise**(`int`):用于提升精度,数据类型支持`int64`,默认值为0。
40 40 
41 > [!NOTE] 41 > [!NOTE]
42 > 当前0、1为保留配置值,2为使能无效行计算,其功能是避免在计算过程中存在整行mask进而导致精度有损失,但是该配置会导致性能下降。42 > 当前0、1为保留配置值,2为使能无效行计算,其功能是避免在计算过程中存在整行mask进而导致精度有损失,但是该配置会导致性能下降。
43 >如果算子可判断出存在无效行场景,会自动使能无效行计算,例如`sparse_mode`为3,Sq \> Skv场景。43 >如果算子可判断出存在无效行场景,会自动使能无效行计算,例如`sparse_mode`为3,Sq \> Skv场景。
44 44 
45-- **prefix**(`List[int]`):可选参数,代表prefix稀疏计算场景每个Batch的N值。数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。45+- **prefix**(`List[int]`):可选参数,代表prefix稀疏计算场景每个Batch的N值。数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
46-- **actual\_seq\_qlen**(`Tensor`):可选参数,varlen场景时需要传入此参数,具体为一维数组的cpu Tensor。表示`query`每个S的累加和长度,数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。46+- **actual\_seq\_qlen**(`Tensor`):可选参数,varlen场景时需要传入此参数,具体为一维数组的cpu Tensor。表示`query`每个S的累加和长度,数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
47 47 
48 比如真正的S长度列表为:2 2 2 2 2,则`actual_seq_qlen`传:2 4 6 8 10。48 比如真正的S长度列表为:2 2 2 2 2,则`actual_seq_qlen`传:2 4 6 8 10。
49 49 
50-- **actual\_seq\_kvlen**(`Tensor`):可选参数,varlen场景时需要传入此参数,具体为一维数组的cpu Tensor。表示`key`/`value`每个S的累加和长度。数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。50+- **actual\_seq\_kvlen**(`Tensor`):可选参数,varlen场景时需要传入此参数,具体为一维数组的cpu Tensor。表示`key`/`value`每个S的累加和长度。数据类型支持`int64`,数据格式支持$ND$。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
51 51 
52 比如真正的S长度列表为:2 2 2 2 2,则actual\_seq\_kvlen传:2 4 6 8 10。52 比如真正的S长度列表为:2 2 2 2 2,则actual\_seq\_kvlen传:2 4 6 8 10。
53 53 
54-- **sparse\_mode**(`int`):表示sparse的模式,可选参数,默认值为0。取值如[表1](#zh-cn_topic_0000001742717129_table1946917414436)所示,不同模式的原理参见[参考资源](#zh-cn_topic_0000001742717129_section28169228374)。当整网的`atten_mask`都相同且shape小于2048\*2048时,建议使用defaultMask模式,来减少内存使用量。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。54+- **sparse\_mode**(`int`):表示sparse的模式,可选参数,默认值为0。取值如[表1](#zh-cn_topic_0000001742717129_table1946917414436)所示,不同模式的原理参见[参考资源](#zh-cn_topic_0000001742717129_section28169228374)。当整网的`atten_mask`都相同且shape小于2048\*2048时,建议使用defaultMask模式,来减少内存使用量。综合约束请见[约束说明](#zh-cn_topic_0000001742717129_section12345537164214)。
55 55 
56 **表1** sparse\_mode不同取值场景说明56 **表1** sparse\_mode不同取值场景说明
57 57 
@@ -130,52 +130,52 @@ torch_npu.npu_fusion_attention_v3(query, key, value, head_num, input_layout, pse
130 </tbody>130 </tbody>
131 </table>131 </table>
132 132 
133-- **gen\_mask\_parallel**(`bool`):DSA生成dropout随机数向量mask的控制开关。默认值为True:同AI Core并行计算;设为False:同AI Core串行计算。133+- **gen\_mask\_parallel**(`bool`):DSA生成dropout随机数向量mask的控制开关。默认值为True:同AI Core并行计算;设为False:同AI Core串行计算。
134-- **sync**(`bool`):DSA生成dropout随机数向量mask的控制开关。默认值为False:dropout mask异步生成;设为True:dropout mask同步生成。134+- **sync**(`bool`):DSA生成dropout随机数向量mask的控制开关。默认值为False:dropout mask异步生成;设为True:dropout mask同步生成。
135-- **softmax_layout**(`string`):可选参数,用于控制TND场景下softmax的输出(softmax_max和softmax_sum)的数据排布方式。当前仅在input_layout=“TND”时进行配置,仅支持传入“TND”。默认情况下,softmax的输出排布为NTD排布;传入TND时,softmax的输出排布为TND排布。135+- **softmax_layout**(`string`):可选参数,用于控制TND场景下softmax的输出(softmax_max和softmax_sum)的数据排布方式。当前仅在input_layout=“TND”时进行配置,仅支持传入“TND”。默认情况下,softmax的输出排布为NTD排布;传入TND时,softmax的输出排布为TND排布。
136-- **sink**(`Tensor`):可选参数,每个注意力头的偏置。shape为`[head_num]`,数据类型仅支持`float32`。136+- **sink**(`Tensor`):可选参数,每个注意力头的偏置。shape为`[head_num]`,数据类型仅支持`float32`。
137 137 
138## 输出说明<a name="zh-cn_topic_0000001742717129_section22231435517"></a>138## 输出说明<a name="zh-cn_topic_0000001742717129_section22231435517"></a>
139 139 
140共6个输出,类型依次为**Tensor、Tensor、Tensor、Tensor、Tensor、Tensor。**140共6个输出,类型依次为**Tensor、Tensor、Tensor、Tensor、Tensor、Tensor。**
141 141 
142-- 第1个输出为`Tensor`,计算公式的最终输出$attention\_out$,数据类型支持`float16`、`bfloat16`、`float32`。142+- 第1个输出为`Tensor`,计算公式的最终输出$attention\_out$,数据类型支持`float16`、`bfloat16`、`float32`。
143-- 第2个输出为`Tensor`,Softmax计算的Max中间结果,用于反向计算,数据类型支持`float`。143+- 第2个输出为`Tensor`,Softmax计算的Max中间结果,用于反向计算,数据类型支持`float`。
144-- 第3个输出为`Tensor`,Softmax计算的Sum中间结果,用于反向计算,数据类型支持`float`。144+- 第3个输出为`Tensor`,Softmax计算的Sum中间结果,用于反向计算,数据类型支持`float`。
145-- 第4个输出为`Tensor`,预留参数,暂未使用。145+- 第4个输出为`Tensor`,预留参数,暂未使用。
146-- 第5个输出为`Tensor`,DSA生成dropout mask中,Philox算法的seed。在aclgraph场景下,返回的是npu Tensor,在非aclgraph场景下,返回的是cpu Tensor。146+- 第5个输出为`Tensor`,DSA生成dropout mask中,Philox算法的seed。在aclgraph场景下,返回的是npu Tensor,在非aclgraph场景下,返回的是cpu Tensor。
147-- 第6个输出为`Tensor`,DSA生成dropout mask中,Philox算法的offset。在aclgraph场景下,返回的是npu Tensor,在非aclgraph场景下,返回的是cpu Tensor。147+- 第6个输出为`Tensor`,DSA生成dropout mask中,Philox算法的offset。在aclgraph场景下,返回的是npu Tensor,在非aclgraph场景下,返回的是cpu Tensor。
148 148 
149## 约束说明<a name="zh-cn_topic_0000001742717129_section12345537164214"></a>149## 约束说明<a name="zh-cn_topic_0000001742717129_section12345537164214"></a>
150 150 
151-- 该接口仅在训练场景下使用。151+- 该接口仅在训练场景下使用。
152-- 输入`query`、`key`、`value`、`pse`的数据类型必须一致。152+- 输入`query`、`key`、`value`、`pse`的数据类型必须一致。
153-- 输入`query`、`key`、`value`的`input_layout`必须一致。153+- 输入`query`、`key`、`value`的`input_layout`必须一致。
154-- 输入`query`、`key`、`value`的shape说明:154+- 输入`query`、`key`、`value`的shape说明:
155- - 输入`key`和`value`的shape必须一致。155+ - 输入`key`和`value`的shape必须一致。
156- - B:batchsize必须相等;非varlen场景B取值范围1\~2M;varlen场景B取值范围1\~2K。156+ - B:batchsize必须相等;非varlen场景B取值范围1\~2M;varlen场景B取值范围1\~2K。
157- - D:Head Dim必须满足Dq=Dk和Dk≥Dv,取值范围1\~768。157+ - D:Head Dim必须满足Dq=Dk和Dk≥Dv,取值范围1\~768。
158- - S:sequence length,取值范围1\~1M。158+ - S:sequence length,取值范围1\~1M。
159 159 
160-- varlen场景下:160+- varlen场景下:
161- - 要求T(B\*S)取值范围1\~1M。161+ - 要求T(B\*S)取值范围1\~1M。
162- - `atten_mask`输入不支持补pad,即`atten_mask`中不能存在某一行全1的场景。162+ - `atten_mask`输入不支持补pad,即`atten_mask`中不能存在某一行全1的场景。
163 163 
164-- 支持输入`query`的N和`key`/`value`的N不相等,但必须成比例关系,即Nq/Nkv必须是非0整数,Nq取值范围1\~256。当Nq/Nkv \> 1时,即为GQA\(grouped-query attention\);当Nq/Nkv=1时,即为MHA\(multi-head attention\)。164+- 支持输入`query`的N和`key`/`value`的N不相等,但必须成比例关系,即Nq/Nkv必须是非0整数,Nq取值范围1\~256。当Nq/Nkv \> 1时,即为GQA\(grouped-query attention\);当Nq/Nkv=1时,即为MHA\(multi-head attention\)。
165 165 
166 > [!NOTE] 166 > [!NOTE]
167 > 本文如无特殊说明,N表示的是Nq。167 > 本文如无特殊说明,N表示的是Nq。
168 168 
169-- `sparse_mode`取值说明:169+- `sparse_mode`取值说明:
170- - `sparse_mode`为1、2、3、4、5、6、7、8时,应传入对应正确的`atten_mask`,否则将导致计算结果错误。当`atten_mask`输入为None时,`sparse_mode`,`pre_tockens`,`next_tockens`参数不生效,固定为全计算。170+ - `sparse_mode`为1、2、3、4、5、6、7、8时,应传入对应正确的`atten_mask`,否则将导致计算结果错误。当`atten_mask`输入为None时,`sparse_mode`,`pre_tockens`,`next_tockens`参数不生效,固定为全计算。
171- - `sparse_mode`配置为1、2、3、5、6时,用户配置的`pre_tockens`、`next_tockens`不会生效。171+ - `sparse_mode`配置为1、2、3、5、6时,用户配置的`pre_tockens`、`next_tockens`不会生效。
172- - `sparse_mode`配置为0、4时,需保证`atten_mask`与`pre_tockens`、`next_tockens`的范围一致。172+ - `sparse_mode`配置为0、4时,需保证`atten_mask`与`pre_tockens`、`next_tockens`的范围一致。
173- - `sparse_mode`配置为7或者8时,不支持可选参数`pse`。173+ - `sparse_mode`配置为7或者8时,不支持可选参数`pse`。
174 174 
175-- `prefix`稀疏计算场景B不大于32,varlen场景不支持非压缩prefix,即不支持sparse\_mode=5;当Sq\>Skv时,`prefix`的N值取值范围\[0, Skv\],当Sq<=Skv时,`prefix`的N值取值范围\[Skv-Sq, Skv\]。175+- `prefix`稀疏计算场景B不大于32,varlen场景不支持非压缩prefix,即不支持sparse\_mode=5;当Sq\>Skv时,`prefix`的N值取值范围\[0, Skv\],当Sq<=Skv时,`prefix`的N值取值范围\[Skv-Sq, Skv\]。
176-- 支持`actual_seq_qlen`中某个Batch上的S长度为0;如果存在S为0的情况,不支持`pse`输入,假设真实的S长度为\[2, 2, 0, 2, 2\],则传入的`actual_seq_qlen`为\[2, 4, 4, 6, 8\]。`actual_seq_qlen`的长度取值范围为1\~2K,varlen场景下长度最大支持1K。176+- 支持`actual_seq_qlen`中某个Batch上的S长度为0;如果存在S为0的情况,不支持`pse`输入,假设真实的S长度为\[2, 2, 0, 2, 2\],则传入的`actual_seq_qlen`为\[2, 4, 4, 6, 8\]。`actual_seq_qlen`的长度取值范围为1\~2K,varlen场景下长度最大支持1K。
177-- TND格式下,支持尾部部分Batch不参与计算,此时`actual_seq_qlen`和`actual_seq_kv_len`尾部传入对应个数个0即可。假设真实的S长度为\[2, 3, 4, 5, 6\],此时后两个Batch不参与计算,则传入的`actual_seq_qlen`为\[2, 5, 9, 0, 0\]。177+- TND格式下,支持尾部部分Batch不参与计算,此时`actual_seq_qlen`和`actual_seq_kv_len`尾部传入对应个数个0即可。假设真实的S长度为\[2, 3, 4, 5, 6\],此时后两个Batch不参与计算,则传入的`actual_seq_qlen`为\[2, 5, 9, 0, 0\]。
178-- 部分场景下,如果计算量过大可能会导致算子执行超时\(aicore error类型报错,errorStr为:timeout or trap error\),此时建议做轴切分处理,注:这里的计算量会受B、S、N、D等参数的影响,值越大计算量越大。178+- 部分场景下,如果计算量过大可能会导致算子执行超时\(aicore error类型报错,errorStr为:timeout or trap error\),此时建议做轴切分处理,注:这里的计算量会受B、S、N、D等参数的影响,值越大计算量越大。
179 179 
180## 调用示例<a name="zh-cn_topic_0000001742717129_section14459801435"></a>180## 调用示例<a name="zh-cn_topic_0000001742717129_section14459801435"></a>
181 181 
@@ -285,12 +285,12 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
285 285 
286**说明:下图中的蓝色表示保留该值,`atten_mask`中,应该配置为False;阴影表示遮蔽该值,`atten_mask`中应配置为True。**286**说明:下图中的蓝色表示保留该值,`atten_mask`中,应该配置为False;阴影表示遮蔽该值,`atten_mask`中应配置为True。**
287 287 
288-- 当`sparse_mode`为0时,代表defaultMask模式。288+- 当`sparse_mode`为0时,代表defaultMask模式。
289- - 不传mask:如果`atten_mask`未传入则不做mask操作,`atten_mask`取值为None,忽略`pre_tockens`和`next_tockens`取值。Masked QK<sup>T</sup>矩阵示意如下:289+ - 不传mask:如果`atten_mask`未传入则不做mask操作,`atten_mask`取值为None,忽略`pre_tockens`和`next_tockens`取值。Masked QK<sup>T</sup>矩阵示意如下:
290 290 
291 ![](../../figures/mode0.png)291 ![](../../figures/mode0.png)
292 292 
293- - `next_tockens`取值为0,`pre_tockens`大于等于Sq,表示causal场景sparse,`atten_mask`应传入下三角矩阵,此时`pre_tockens`和`next_tockens`之间的部分需要计算,Masked QK<sup>T</sup>矩阵示意如下:293+ - `next_tockens`取值为0,`pre_tockens`大于等于Sq,表示causal场景sparse,`atten_mask`应传入下三角矩阵,此时`pre_tockens`和`next_tockens`之间的部分需要计算,Masked QK<sup>T</sup>矩阵示意如下:
294 294 
295 ![](../../figures/1.png)295 ![](../../figures/1.png)
296 296 
@@ -298,7 +298,7 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
298 298 
299 ![](../../figures/1-1.png)299 ![](../../figures/1-1.png)
300 300 
301- - `pre_tockens`小于Sq,`next_tockens`小于Skv,且都大于等于0,表示band场景,此时`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>矩阵示意如下:301+ - `pre_tockens`小于Sq,`next_tockens`小于Skv,且都大于等于0,表示band场景,此时`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>矩阵示意如下:
302 302 
303 ![](../../figures/1-2.png)303 ![](../../figures/1-2.png)
304 304 
@@ -306,25 +306,25 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
306 306 
307 ![](../../figures/1-3.png)307 ![](../../figures/1-3.png)
308 308 
309- - `next_tockens`为负数,以pre\_tockens=9,next\_tockens=-3为例,`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>示意如下:309+ - `next_tockens`为负数,以pre\_tockens=9,next\_tockens=-3为例,`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>示意如下:
310 310 
311 **说明:`next_tockens`为负数时,`pre_tockens`取值必须大于等于`next_tockens`的绝对值,且`next_tockens`的绝对值小于Skv。**311 **说明:`next_tockens`为负数时,`pre_tockens`取值必须大于等于`next_tockens`的绝对值,且`next_tockens`的绝对值小于Skv。**
312 312 
313 ![](../../figures/1-4.png)313 ![](../../figures/1-4.png)
314 314 
315- - `pre_tockens`为负数,以next\_tockens=7,pre\_tockens=-3为例,`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>示意如下:315+ - `pre_tockens`为负数,以next\_tockens=7,pre\_tockens=-3为例,`pre_tockens`和`next_tockens`之间的部分需要计算。Masked QK<sup>T</sup>示意如下:
316 316 
317 **说明:`pre_tockens`为负数时,`next_tockens`取值必须大于等于`pre_tockens`的绝对值,且`pre_tockens`的绝对值小于Sq。**317 **说明:`pre_tockens`为负数时,`next_tockens`取值必须大于等于`pre_tockens`的绝对值,且`pre_tockens`的绝对值小于Sq。**
318 318 
319 ![](../../figures/1-5.png)319 ![](../../figures/1-5.png)
320 320 
321-- 当`sparse_mode`为1时,代表allMask,即传入完整的`atten_mask`矩阵。321+- 当`sparse_mode`为1时,代表allMask,即传入完整的`atten_mask`矩阵。
322 322 
323 该场景下忽略`next_tockens`、`pre_tockens`取值,Masked QK<sup>T</sup>矩阵示意如下:323 该场景下忽略`next_tockens`、`pre_tockens`取值,Masked QK<sup>T</sup>矩阵示意如下:
324 324 
325 ![](../../figures/1-6.png)325 ![](../../figures/1-6.png)
326 326 
327-- 当`sparse_mode`为2时,代表leftUpCausal模式的mask,对应以左上顶点划分的下三角场景(参数起点为左上角)。该场景下忽略`pre_tockens`、`next_tockens`取值,Masked QK<sup>T</sup>矩阵示意如下:327+- 当`sparse_mode`为2时,代表leftUpCausal模式的mask,对应以左上顶点划分的下三角场景(参数起点为左上角)。该场景下忽略`pre_tockens`、`next_tockens`取值,Masked QK<sup>T</sup>矩阵示意如下:
328 328 
329 ![](../../figures/1-7.png)329 ![](../../figures/1-7.png)
330 330 
@@ -332,15 +332,15 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
332 332 
333 ![](../../figures/1-8.png)333 ![](../../figures/1-8.png)
334 334 
335-- 当`sparse_mode`为3时,代表rightDownCausal模式的mask,对应以右下顶点划分的下三角场景(参数起点为右下角)。该场景下忽略`pre_tockens`、`next_tockens`取值。`atten_mask`为优化后的压缩下三角矩阵(2048\*2048),Masked QK<sup>T</sup>矩阵示意如下:335+- 当`sparse_mode`为3时,代表rightDownCausal模式的mask,对应以右下顶点划分的下三角场景(参数起点为右下角)。该场景下忽略`pre_tockens`、`next_tockens`取值。`atten_mask`为优化后的压缩下三角矩阵(2048\*2048),Masked QK<sup>T</sup>矩阵示意如下:
336 336 
337 ![](../../figures/1-9.png)337 ![](../../figures/1-9.png)
338 338 
339-- 当`sparse_mode`为4时,代表band场景,即计算`pre_tockens`和`next_tockens`之间的部分,参数起点为右下角,`pre_tockens`和`next_tockens`之间需要有交集。`atten_mask`为优化后的压缩下三角矩阵(2048\*2048)。Masked QK<sup>T</sup>矩阵示意如下:339+- 当`sparse_mode`为4时,代表band场景,即计算`pre_tockens`和`next_tockens`之间的部分,参数起点为右下角,`pre_tockens`和`next_tockens`之间需要有交集。`atten_mask`为优化后的压缩下三角矩阵(2048\*2048)。Masked QK<sup>T</sup>矩阵示意如下:
340 340 
341 ![](../../figures/1-10.png)341 ![](../../figures/1-10.png)
342 342 
343-- 当`sparse_mode`为5时,代表prefix非压缩场景,即在rightDownCausal的基础上,左侧加上一个长为Sq,宽为N的矩阵,N的值由可选参数prefix获取,例如下图中表示batch=2场景下prefix传入数组\[4,5\],每个batch轴的N值可以不一样,参数起点为左上角。343+- 当`sparse_mode`为5时,代表prefix非压缩场景,即在rightDownCausal的基础上,左侧加上一个长为Sq,宽为N的矩阵,N的值由可选参数prefix获取,例如下图中表示batch=2场景下prefix传入数组\[4,5\],每个batch轴的N值可以不一样,参数起点为左上角。
344 344 
345 该场景下忽略`pre_tockens`、`next_tockens`取值,`atten_mask`矩阵数据格式须为BNSS或B1SS,Masked QK<sup>T</sup>矩阵示意如下:345 该场景下忽略`pre_tockens`、`next_tockens`取值,`atten_mask`矩阵数据格式须为BNSS或B1SS,Masked QK<sup>T</sup>矩阵示意如下:
346 346 
@@ -350,39 +350,39 @@ QK<sup>T</sup>矩阵在`atten_mask`为True的位置会被遮蔽,效果如下
350 350 
351 ![](../../figures/1-12.png)351 ![](../../figures/1-12.png)
352 352 
353-- 当`sparse_mode`为6时,代表prefix压缩场景,即prefix场景时,attenMask为优化后的压缩下三角+矩形的矩阵(3072\*2048):其中上半部分\[2048,2048\]的下三角矩阵,下半部分为\[1024,2048\]的矩形矩阵,矩形矩阵左半部分全0,右半部分全1,`atten_mask`应传入矩阵示意如下。该场景下忽略`pre_tockens`、`next_tockens`取值。353+- 当`sparse_mode`为6时,代表prefix压缩场景,即prefix场景时,attenMask为优化后的压缩下三角+矩形的矩阵(3072\*2048):其中上半部分\[2048,2048\]的下三角矩阵,下半部分为\[1024,2048\]的矩形矩阵,矩形矩阵左半部分全0,右半部分全1,`atten_mask`应传入矩阵示意如下。该场景下忽略`pre_tockens`、`next_tockens`取值。
354 354 
355 ![](../../figures/1-13.png)355 ![](../../figures/1-13.png)
356 356 
357-- 当`sparse_mode`为7时,表示varlen且为长序列外切场景(即长序列在模型脚本中进行多卡切query的sequence length);用户需要确保外切前为使用sparse\_mode=3的场景;当前mode下用户需要设置`pre_tockens`和`next_tockens`(起点为右下顶点),且需要保证参数正确,否则会存在精度问题。357+- 当`sparse_mode`为7时,表示varlen且为长序列外切场景(即长序列在模型脚本中进行多卡切query的sequence length);用户需要确保外切前为使用sparse\_mode=3的场景;当前mode下用户需要设置`pre_tockens`和`next_tockens`(起点为右下顶点),且需要保证参数正确,否则会存在精度问题。
358 358 
359 Masked QK<sup>T</sup>矩阵示意如下,在第二个batch对`query`进行切分,`key`和`value`不切分,4x6的mask矩阵被切分成2x6和2x6的mask,分别在卡1和卡2上计算:359 Masked QK<sup>T</sup>矩阵示意如下,在第二个batch对`query`进行切分,`key`和`value`不切分,4x6的mask矩阵被切分成2x6和2x6的mask,分别在卡1和卡2上计算:
360 360 
361- - 卡1的最后一块mask为band类型的mask,配置pre\_tockens=6(保证大于等于最后一个Skv),next\_tockens=-2,`actual_seq_qlen`应传入\{3,5\},`actual_seq_kvlen`应传入\{3,9\}。361+ - 卡1的最后一块mask为band类型的mask,配置pre\_tockens=6(保证大于等于最后一个Skv),next\_tockens=-2,`actual_seq_qlen`应传入\{3,5\},`actual_seq_kvlen`应传入\{3,9\}。
362- - 卡2的mask类型切分后不变,`sparse_mode`为3,`actual_seq_qlen`应传入\{2,7,11\},`actual_seq_kvlen`应传入\{6,11,15\}。362+ - 卡2的mask类型切分后不变,`sparse_mode`为3,`actual_seq_qlen`应传入\{2,7,11\},`actual_seq_kvlen`应传入\{6,11,15\}。
363 363 
364 ![](../../figures/1-14.png)364 ![](../../figures/1-14.png)
365 365 
366 > [!NOTE] 366 > [!NOTE]
367- > - 如果配置sparse\_mode=7,但实际只存在一个batch,用户需按照band模式的要求来配置参数;sparse\_mode=7时,用户需要输入2048x2048的下三角mask作为该融合算子的输入。367+ > - 如果配置sparse\_mode=7,但实际只存在一个batch,用户需按照band模式的要求来配置参数;sparse\_mode=7时,用户需要输入2048x2048的下三角mask作为该融合算子的输入。
368- > - 基于sparse\_mode=3进行外切产生的band模式的sparse的参数应符合以下条件:368+ > - 基于sparse\_mode=3进行外切产生的band模式的sparse的参数应符合以下条件:
369- > - pre\_tockens \>= last\_Skv。369+ > - pre\_tockens \>= last\_Skv。
370- > - next\_tockens <= 0。370+ > - next\_tockens <= 0。
371- > - 当前模式下不支持可选输入pse。371+ > - 当前模式下不支持可选输入pse。
372 372 
373-- 当`sparse_mode`为8时,表示varlen且为长序列外切场景;用户需要确保外切前为使用sparse\_mode=2的场景;当前mode下用户需要设置`pre_tockens`和`next_tockens`(起点为右下顶点),且需要保证参数正确,否则会存在精度问题。373+- 当`sparse_mode`为8时,表示varlen且为长序列外切场景;用户需要确保外切前为使用sparse\_mode=2的场景;当前mode下用户需要设置`pre_tockens`和`next_tockens`(起点为右下顶点),且需要保证参数正确,否则会存在精度问题。
374 374 
375 Masked QK<sup>T</sup>矩阵示意如下,在第二个batch对`query`进行切分,`key`和`value`不切分,5x4的mask矩阵被切分成2x4和3x4的mask,分别在卡1和卡2上计算:375 Masked QK<sup>T</sup>矩阵示意如下,在第二个batch对`query`进行切分,`key`和`value`不切分,5x4的mask矩阵被切分成2x4和3x4的mask,分别在卡1和卡2上计算:
376 376 
377- - 卡1的mask类型切分后不变,`sparse_mode`为2,`actual_seq_qlen`应传入\{3,5\},`actual_seq_kvlen`应传入\{3,7\}。377+ - 卡1的mask类型切分后不变,`sparse_mode`为2,`actual_seq_qlen`应传入\{3,5\},`actual_seq_kvlen`应传入\{3,7\}。
378- - 卡2的第一块mask为band类型的mask,配置pre\_tockens=4(保证大于等于第一个Skv),next\_tockens=1,`actual_seq_qlen`应传入\{3,8,12\},`actual_seq_kvlen`应传入\{4,9,13\}。378+ - 卡2的第一块mask为band类型的mask,配置pre\_tockens=4(保证大于等于第一个Skv),next\_tockens=1,`actual_seq_qlen`应传入\{3,8,12\},`actual_seq_kvlen`应传入\{4,9,13\}。
379 379 
380 ![](../../figures/1-15.png)380 ![](../../figures/1-15.png)
381 381 
382 > [!NOTE] 382 > [!NOTE]
383- > - 如果配置sparse\_mode=8,但实际只存在一个batch,用户需按照band模式的要求来配置参数;sparse\_mode=8时,用户需要输入2048x2048的下三角mask作为该融合算子的输入。383+ > - 如果配置sparse\_mode=8,但实际只存在一个batch,用户需按照band模式的要求来配置参数;sparse\_mode=8时,用户需要输入2048x2048的下三角mask作为该融合算子的输入。
384- > - 基于sparse\_mode=2进行外切产生的band模式的sparse的参数应符合以下条件:384+ > - 基于sparse\_mode=2进行外切产生的band模式的sparse的参数应符合以下条件:
385- > - pre\_tockens \>= first\_Skv。385+ > - pre\_tockens \>= first\_Skv。
386- > - next\_tockens范围无约束,根据实际情况进行配置。386+ > - next\_tockens范围无约束,根据实际情况进行配置。
387- > - 当前模式下不支持可选输入pse。387+ > - 当前模式下不支持可选输入pse。
388- 388+
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_gather_sparse_index.md+4-7
@@ -7,7 +7,7 @@
7|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9 9 
10-## 功能说明: 10+## 功能说明
11 11 
12- API功能:从输入Tensor的指定维度,按照`index`中的下标序号提取元素,保存到输出Tensor中。12- API功能:从输入Tensor的指定维度,按照`index`中的下标序号提取元素,保存到输出Tensor中。
13 13 
@@ -40,12 +40,9 @@
40 \end{bmatrix}40 \end{bmatrix}
41 $$41 $$
42 42 
43- 
44- 
45## 函数原型43## 函数原型
46 44 
47- 45+```python
48-```
49torch_npu.npu_gather_sparse_index(input, index) -> Tensor46torch_npu.npu_gather_sparse_index(input, index) -> Tensor
50```47```
51 48 
@@ -61,13 +58,13 @@ torch_npu.npu_gather_sparse_index(input, index) -> Tensor
61接口计算获得的结果,包含按照`index`中的下标序号提取的元素。数据类型与`input`一致,输出维度为$index.dim + input.dim - 1$。例如`input.shape = [16, 32]`, `index.shape = [2, 3]`,则输出张量 `out.shape = [2, 3, 32]`58接口计算获得的结果,包含按照`index`中的下标序号提取的元素。数据类型与`input`一致,输出维度为$index.dim + input.dim - 1$。例如`input.shape = [16, 32]`, `index.shape = [2, 3]`,则输出张量 `out.shape = [2, 3, 32]`
62 59 
63## 约束说明60## 约束说明
61+ 
64- `input`的维度与`index`的维度之和减1不能超过8,即$index.dim + input.dim - 1<=8$。62- `input`的维度与`index`的维度之和减1不能超过8,即$index.dim + input.dim - 1<=8$。
65- 为获取性能收益,`input``index`需要满足如下约束:63- 为获取性能收益,`input``index`需要满足如下约束:
66 1. `input`的shape内积需要大于$150 * 1024 / itemsize$,其中itemsize为`input` dtype对应元素大小,可以通过`torch.dtype.itemsize`查询。64 1. `input`的shape内积需要大于$150 * 1024 / itemsize$,其中itemsize为`input` dtype对应元素大小,可以通过`torch.dtype.itemsize`查询。
67 2. `index`的shape内积大于960。65 2. `index`的shape内积大于960。
68 3. 数据需要聚合,即非0值分布集中,0值分布集中。66 3. 数据需要聚合,即非0值分布集中,0值分布集中。
69 67 
70- 
71## 调用示例68## 调用示例
72 69 
73```python70```python
@@ -77,4 +74,4 @@ import torch_npu
77inputs = torch.randn(16, 32).npu()74inputs = torch.randn(16, 32).npu()
78index = torch.randint(0, 16, [2, 3]).npu()75index = torch.randint(0, 16, [2, 3]).npu()
79out = torch_npu.npu_gather_sparse_index(inputs, index)76out = torch_npu.npu_gather_sparse_index(inputs, index)
80-```77+```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_gelu.md+2-3
@@ -29,7 +29,7 @@
29 29 
30## 函数原型30## 函数原型
31 31 
32-```32+```python
33torch_npu.npu_gelu(input, approximate='none') -> Tensor33torch_npu.npu_gelu(input, approximate='none') -> Tensor
34```34```
35 35 
@@ -43,11 +43,11 @@ torch_npu.npu_gelu(input, approximate='none') -> Tensor
43- **approximate** (`String`):可选参数,字符串类型,计算使用的激活函数模式,可配置为`none`或者`tanh`。其中`none`代表使用erf模式,`tanh`代表使用tanh模式。43- **approximate** (`String`):可选参数,字符串类型,计算使用的激活函数模式,可配置为`none`或者`tanh`。其中`none`代表使用erf模式,`tanh`代表使用tanh模式。
44 44 
45## 返回值说明45## 返回值说明
46+ 
46`Tensor`47`Tensor`
47 48 
48数据类型必须和`input`一样,数据格式支持$ND$,shape必须和`input`一样,支持非连续的Tensor。49数据类型必须和`input`一样,数据格式支持$ND$,shape必须和`input`一样,支持非连续的Tensor。
49 50 
50- 
51## 约束说明51## 约束说明
52 52 
53- 该接口支持图模式。53- 该接口支持图模式。
@@ -112,4 +112,3 @@ torch_npu.npu_gelu(input, approximate='none') -> Tensor
112 # 执行上述代码的输出类似如下112 # 执行上述代码的输出类似如下
113 torch.Size([100, 10, 20]) torch.float32113 torch.Size([100, 10, 20]) torch.float32
114 ```114 ```
115- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_gelu_mul.md+3-3
@@ -39,10 +39,9 @@
39 39
40 其中$\text{out}$形状与原始输入`input`完全一致。40 其中$\text{out}$形状与原始输入`input`完全一致。
41 41 
42- 
43## 函数原型42## 函数原型
44 43 
45-```44+```python
46torch_npu.npu_gelu_mul(input, *, approximate="none") -> Tensor45torch_npu.npu_gelu_mul(input, *, approximate="none") -> Tensor
47```46```
48 47 
@@ -54,6 +53,7 @@ torch_npu.npu_gelu_mul(input, *, approximate="none") -> Tensor
54 - "tanh":使用双曲正切(tanh)近似模式,计算效率高,适用于大规模训练或推理加速场景。53 - "tanh":使用双曲正切(tanh)近似模式,计算效率高,适用于大规模训练或推理加速场景。
55 54 
56## 返回值说明55## 返回值说明
56+ 
57`Tensor`57`Tensor`
58 58 
59输出张量,对应公式中的$out$,数据类型支持bfloat16、float16、float。shape维度2至8维。支持非连续的Tensor,数据格式支持$ND$,输出的数据类型与输入`input`保持一致,输出shape和输入shape其他维度一致,最后一维的值为输入shape最后一维值的二分之一。59输出张量,对应公式中的$out$,数据类型支持bfloat16、float16、float。shape维度2至8维。支持非连续的Tensor,数据格式支持$ND$,输出的数据类型与输入`input`保持一致,输出shape和输入shape其他维度一致,最后一维的值为输入shape最后一维值的二分之一。
@@ -67,4 +67,4 @@ torch_npu.npu_gelu_mul(input, *, approximate="none") -> Tensor
67>>> mode = "tanh"67>>> mode = "tanh"
68>>> output = torch_npu.npu_gelu_mul(input, approximate=mode)68>>> output = torch_npu.npu_gelu_mul(input, approximate=mode)
69 69 
70-```70+```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_gmm_alltoallv.md+42-43
@@ -19,72 +19,72 @@
19 19 
20## 功能说明<a name="zh-cn_topic_0000002317314449_section14441124184110"></a>20## 功能说明<a name="zh-cn_topic_0000002317314449_section14441124184110"></a>
21 21 
22-- API功能:MoE网络中,完成路由专家GroupedMatMul、AlltoAllv融合并实现与共享专家MatMul并行融合,先计算后通信。22+- API功能:MoE网络中,完成路由专家GroupedMatMul、AlltoAllv融合并实现与共享专家MatMul并行融合,先计算后通信。
23-- 路由专家计算公式:23+- 路由专家计算公式:
24 24 
25 ![](../../figures/zh-cn_formulaimage_0000002323688460.png)25 ![](../../figures/zh-cn_formulaimage_0000002323688460.png)
26 26 
27- - gmm\_x指路由专家GroupedMatMul计算的左矩阵。27+ - gmm\_x指路由专家GroupedMatMul计算的左矩阵。
28- - gmm\_weight指路由专家GroupedMatMul计算的右矩阵。28+ - gmm\_weight指路由专家GroupedMatMul计算的右矩阵。
29- - gmm\_y指路由专家进行GroupedMatMul计算的输出,后续用于Unpermute计算。29+ - gmm\_y指路由专家进行GroupedMatMul计算的输出,后续用于Unpermute计算。
30- - unpermute\_out是gmm\_y进行Unpermute计算的输出结果,作为AlltoAllv通信的输入。30+ - unpermute\_out是gmm\_y进行Unpermute计算的输出结果,作为AlltoAllv通信的输入。
31- - y指对unpermute\_out进行AlltoAllv通信输出。31+ - y指对unpermute\_out进行AlltoAllv通信输出。
32 32 
33-- 共享专家计算公式:33+- 共享专家计算公式:
34 34 
35 ![](../../figures/zh-cn_formulaimage_0000002323838248.png)35 ![](../../figures/zh-cn_formulaimage_0000002323838248.png)
36 36 
37- - mm\_x指共享专家MatMul计算的左矩阵。37+ - mm\_x指共享专家MatMul计算的左矩阵。
38- - mm\_weight指共享专家MatMul计算的右矩阵。38+ - mm\_weight指共享专家MatMul计算的右矩阵。
39- - mm\_y指共享专家MatMul计算的输出。39+ - mm\_y指共享专家MatMul计算的输出。
40 40 
41## 函数原型<a name="zh-cn_topic_0000002317314449_section45077510411"></a>41## 函数原型<a name="zh-cn_topic_0000002317314449_section45077510411"></a>
42 42 
43-```43+```python
44torch_npu.npu_gmm_alltoallv(gmm_x, gmm_weight, hcom, ep_world_size, send_counts, recv_counts, *, send_counts_tensor=None, recv_counts_tensor=None, mm_x=None, mm_weight=None, trans_gmm_weight=False, trans_mm_weight=False) -> (Tensor, Tensor)44torch_npu.npu_gmm_alltoallv(gmm_x, gmm_weight, hcom, ep_world_size, send_counts, recv_counts, *, send_counts_tensor=None, recv_counts_tensor=None, mm_x=None, mm_weight=None, trans_gmm_weight=False, trans_mm_weight=False) -> (Tensor, Tensor)
45```45```
46 46 
47## 参数说明<a name="zh-cn_topic_0000002317314449_section112637109429"></a>47## 参数说明<a name="zh-cn_topic_0000002317314449_section112637109429"></a>
48 48 
49-- **gmm\_x**(`Tensor`):必选参数,GroupedMatMul计算的左矩阵。数据类型支持`float16`、`bfloat16`,支持2维,shape为$(A, H1)$,数据格式支持ND。49+- **gmm\_x**(`Tensor`):必选参数,GroupedMatMul计算的左矩阵。数据类型支持`float16`、`bfloat16`,支持2维,shape为$(A, H1)$,数据格式支持ND。
50-- **gmm\_weight**(`Tensor`):必选参数,GroupedMatMul计算的右矩阵。数据类型与`gmm_x`保持一致,支持3维,shape为$(e, H1, N1)$,数据格式支持ND。50+- **gmm\_weight**(`Tensor`):必选参数,GroupedMatMul计算的右矩阵。数据类型与`gmm_x`保持一致,支持3维,shape为$(e, H1, N1)$,数据格式支持ND。
51-- **hcom**(`str`):必选参数,专家并行的通信域名,字符串长度要求\(0, 128\)。51+- **hcom**(`str`):必选参数,专家并行的通信域名,字符串长度要求\(0, 128\)。
52-- **ep\_world\_size**(`int`):必选参数,EP通信域size,取值支持8、16、32、64、128。52+- **ep\_world\_size**(`int`):必选参数,EP通信域size,取值支持8、16、32、64、128。
53-- **send\_counts**(`List[int]`):必选参数,为一个列表,表示发送给其他卡的token数,列表长度为卡数。列表中元素的数据类型支持`int`,取值为e\*`ep_world_size`,最大值为256。53+- **send\_counts**(`List[int]`):必选参数,为一个列表,表示发送给其他卡的token数,列表长度为卡数。列表中元素的数据类型支持`int`,取值为e\*`ep_world_size`,最大值为256。
54-- **recv\_counts**(`List[int]`):必选参数,为一个列表,表示接收其他卡的token数,列表长度为卡数。列表中元素的数据类型支持`int`,取值大小为e\*`ep_world_size`,最大值为256。54+- **recv\_counts**(`List[int]`):必选参数,为一个列表,表示接收其他卡的token数,列表长度为卡数。列表中元素的数据类型支持`int`,取值大小为e\*`ep_world_size`,最大值为256。
55-- **send\_counts\_tensor**(`Tensor`):可选参数,数据类型支持`int`,shape为$(e*ep\_world\_size,)$,数据格式支持ND。**当前版本暂不支持**,使用默认值即可。55+- **send\_counts\_tensor**(`Tensor`):可选参数,数据类型支持`int`,shape为$(e*ep\_world\_size,)$,数据格式支持ND。**当前版本暂不支持**,使用默认值即可。
56-- **recv\_counts\_tensor**(`Tensor`):可选参数,数据类型支持`int`,shape为$(e*ep\_world\_size,)$,数据格式支持ND。**当前版本暂不支持**,使用默认值即可。56+- **recv\_counts\_tensor**(`Tensor`):可选参数,数据类型支持`int`,shape为$(e*ep\_world\_size,)$,数据格式支持ND。**当前版本暂不支持**,使用默认值即可。
57-- **mm\_x**(`Tensor`):可选参数,共享专家MatMul计算中的左矩阵。当需要融合共享专家矩阵计算时,该参数必选,数据类型支持`float16`、`bfloat16`,支持2维,shape为$(BS, H2)$。57+- **mm\_x**(`Tensor`):可选参数,共享专家MatMul计算中的左矩阵。当需要融合共享专家矩阵计算时,该参数必选,数据类型支持`float16`、`bfloat16`,支持2维,shape为$(BS, H2)$。
58-- **mm\_weight**(`Tensor`):可选参数,共享专家MatMul计算中的右矩阵。当需要融合共享专家矩阵计算时,该参数必选,数据类型与`mm_x`保持一致,支持2维,shape为$(H2, N2)$。58+- **mm\_weight**(`Tensor`):可选参数,共享专家MatMul计算中的右矩阵。当需要融合共享专家矩阵计算时,该参数必选,数据类型与`mm_x`保持一致,支持2维,shape为$(H2, N2)$。
59-- **trans\_gmm\_weight**(`bool`):可选参数,GroupedMatMul的右矩阵是否需要转置,true表示需要转置,false表示不转置。59+- **trans\_gmm\_weight**(`bool`):可选参数,GroupedMatMul的右矩阵是否需要转置,true表示需要转置,false表示不转置。
60-- **trans\_mm\_weight**(`bool`):可选参数,共享专家MatMul的右矩阵是否需要转置,true表示需要转置,false表示不转置。60+- **trans\_mm\_weight**(`bool`):可选参数,共享专家MatMul的右矩阵是否需要转置,true表示需要转置,false表示不转置。
61 61 
62## 返回值说明<a name="zh-cn_topic_0000002317314449_section22231435517"></a>62## 返回值说明<a name="zh-cn_topic_0000002317314449_section22231435517"></a>
63 63 
64-- **y**(`Tensor`):表示最终计算结果,数据类型与输入`gmm_x`保持一致,支持2维,shape为$(BSK, N1)$。64+- **y**(`Tensor`):表示最终计算结果,数据类型与输入`gmm_x`保持一致,支持2维,shape为$(BSK, N1)$。
65-- **mm\_y**(`Tensor`):共享专家MatMul的输出,数据类型与`mm_x`保持一致,支持2维,shape为$(BS, N2)$。仅当传入`mm_x`与`mm_weight`才输出。65+- **mm\_y**(`Tensor`):共享专家MatMul的输出,数据类型与`mm_x`保持一致,支持2维,shape为$(BS, N2)$。仅当传入`mm_x`与`mm_weight`才输出。
66 66 
67## 约束说明<a name="zh-cn_topic_0000002317314449_section12345537164214"></a>67## 约束说明<a name="zh-cn_topic_0000002317314449_section12345537164214"></a>
68 68 
69-- 该接口支持推理场景下使用。69+- 该接口支持推理场景下使用。
70-- 该接口支持图模式。70+- 该接口支持图模式。
71-- 单卡通信量取值大于等于2MB。71+- 单卡通信量取值大于等于2MB。
72-- 输入参数Tensor中shape使用的变量说明:72+- 输入参数Tensor中shape使用的变量说明:
73- - BSK:本卡接收的token数(BS\*K=BSK),是recv\_counts参数累加之和,取值范围\(0, 52428800\)。73+ - BSK:本卡接收的token数(BS\*K=BSK),是recv\_counts参数累加之和,取值范围\(0, 52428800\)。
74 74 
75- - H1:表示路由专家hidden size隐藏层大小,取值范围\(0, 65536\)。75+ - H1:表示路由专家hidden size隐藏层大小,取值范围\(0, 65536\)。
76- - H2:表示共享专家hidden size隐藏层大小,取值范围\(0, 12288\]。76+ - H2:表示共享专家hidden size隐藏层大小,取值范围\(0, 12288\]。
77- - e:表示单卡上专家个数,e<=32,e \* ep\_world\_size最大支持256。77+ - e:表示单卡上专家个数,e<=32,e \* ep\_world\_size最大支持256。
78- - N1:表示路由专家的head\_num,取值范围\(0, 65536\)。78+ - N1:表示路由专家的head\_num,取值范围\(0, 65536\)。
79- - N2:表示共享专家的head\_num,取值范围\(0, 65536\)。79+ - N2:表示共享专家的head\_num,取值范围\(0, 65536\)。
80- - BS:batch sequence size。80+ - BS:batch sequence size。
81- - K:表示选取top\_k个专家,K的范围\[2, 8\]。81+ - K:表示选取top\_k个专家,K的范围\[2, 8\]。
82- - A:本卡发送的token数,是send\_counts参数累加之和。82+ - A:本卡发送的token数,是send\_counts参数累加之和。
83- - EP通信域内所有卡上的A参数的累加和等于所有卡上的BSK参数的累加和。83+ - EP通信域内所有卡上的A参数的累加和等于所有卡上的BSK参数的累加和。
84 84 
85## 调用示例<a name="zh-cn_topic_0000002317314449_section14459801435"></a>85## 调用示例<a name="zh-cn_topic_0000002317314449_section14459801435"></a>
86 86 
87-- 单算子模式调用87+- 单算子模式调用
88 88 
89 ```python89 ```python
90 import torch90 import torch
@@ -136,7 +136,7 @@ torch_npu.npu_gmm_alltoallv(gmm_x, gmm_weight, hcom, ep_world_size, send_counts,
136 mp.spawn(run_npu_gmm_alltoallv, args=(epWorkSize, master_ip, master_port, gmm_x_shape, gmm_weight_shape, send_counts, recv_counts, dtype), nprocs=epWorkSize)136 mp.spawn(run_npu_gmm_alltoallv, args=(epWorkSize, master_ip, master_port, gmm_x_shape, gmm_weight_shape, send_counts, recv_counts, dtype), nprocs=epWorkSize)
137 ```137 ```
138 138 
139-- 图模式调用139+- 图模式调用
140 140 
141 ```python141 ```python
142 import torch142 import torch
@@ -213,4 +213,3 @@ torch_npu.npu_gmm_alltoallv(gmm_x, gmm_weight, hcom, ep_world_size, send_counts,
213 213
214 mp.spawn(run_npu_gmm_alltoallv, args=(epWorkSize, master_ip, master_port, gmm_x_shape, gmm_weight_shape, send_counts, recv_counts, dtype), nprocs=epWorkSize)214 mp.spawn(run_npu_gmm_alltoallv, args=(epWorkSize, master_ip, master_port, gmm_x_shape, gmm_weight_shape, send_counts, recv_counts, dtype), nprocs=epWorkSize)
215 ```215 ```
216- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_group_norm_silu.md+36-37
@@ -9,9 +9,9 @@
9 9 
10## 功能说明10## 功能说明
11 11 
12-- API功能:计算输入张量`input`按组归一化的结果,包括张量out、均值meanOut、标准差的倒数rstdOut以及silu的输出。12+- API功能:计算输入张量`input`按组归一化的结果,包括张量out、均值meanOut、标准差的倒数rstdOut以及silu的输出。
13-- 计算公式:13+- 计算公式:
14- - GroupNorm:$x$为输入`input`,$\gamma$和$\beta$分别代表输入`weight`和`bias`,$E[x] = \bar{x}$代表$x$的均值,$ Var[x]=\frac{1}{n}\sum_{i=1}^{n} (x_i - E[x])^2 $ 代表$x$的方差,则14+ - GroupNorm:$x$为输入`input`,$\gamma$和$\beta$分别代表输入`weight`和`bias`,$E[x] = \bar{x}$代表$x$的均值,$ Var[x]=\frac{1}{n}\sum_{i=1}^{n} (x_i - E[x])^2 $ 代表$x$的方差,则
15 $$15 $$
16 \begin{cases}16 \begin{cases}
17 \text{groupnormOut} = \frac{x - E[x]}{\sqrt{Var[x] + eps}} * \gamma + \beta \\17 \text{groupnormOut} = \frac{x - E[x]}{\sqrt{Var[x] + eps}} * \gamma + \beta \\
@@ -19,61 +19,61 @@
19 \text{rstdOut} = \frac{1}{\sqrt{Var[x] + eps}}19 \text{rstdOut} = \frac{1}{\sqrt{Var[x] + eps}}
20 \end{cases}20 \end{cases}
21 $$21 $$
22- - Silu:22+ - Silu:
23 $$23 $$
24 \text{out} = \frac{\text{groupnormOut}}{1 + e^{-\text{groupnormOut}}}24 \text{out} = \frac{\text{groupnormOut}}{1 + e^{-\text{groupnormOut}}}
25 $$25 $$
26 26 
27## 函数原型27## 函数原型
28 28 
29-```29+```python
30torch_npu.npu_group_norm_silu(input, weight, bias, group, eps=0.00001) -> (Tensor, Tensor, Tensor)30torch_npu.npu_group_norm_silu(input, weight, bias, group, eps=0.00001) -> (Tensor, Tensor, Tensor)
31```31```
32 32 
33## 参数说明33## 参数说明
34 34 
35-- **input** (`Tensor`):必选参数,源数据张量,维度需要为2~8维且第1维度能整除`group`。数据格式支持$ND$,支持非连续的Tensor。35+- **input** (`Tensor`):必选参数,源数据张量,维度需要为2~8维且第1维度能整除`group`。数据格式支持$ND$,支持非连续的Tensor。
36- - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。36+ - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。
37- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。37+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。
38 38 
39-- **weight** (`Tensor`):可选参数,索引张量,维度为1且元素数量需与输入`input`的第1维度保持相同,数据格式支持$ND$,支持非连续的Tensor。39+- **weight** (`Tensor`):可选参数,索引张量,维度为1且元素数量需与输入`input`的第1维度保持相同,数据格式支持$ND$,支持非连续的Tensor。
40- - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。40+ - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。
41- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。41+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。
42 42 
43-- **bias** (`Tensor`):可选参数,更新数据张量,维度为1且元素数量需与输入`input`的第1维度保持相同,数据格式支持$ND$,支持非连续的Tensor。43+- **bias** (`Tensor`):可选参数,更新数据张量,维度为1且元素数量需与输入`input`的第1维度保持相同,数据格式支持$ND$,支持非连续的Tensor。
44- - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。44+ - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。
45- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。45+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。
46 46 
47-- **group** (`int`):必选参数,表示将输入`input`的第1维度分为group组,group需大于0。47+- **group** (`int`):必选参数,表示将输入`input`的第1维度分为group组,group需大于0。
48-- **eps** (`float`):可选参数,数值稳定性而加到分母上的值,若保持精度,则eps需大于0。默认值为0.00001。48+- **eps** (`float`):可选参数,数值稳定性而加到分母上的值,若保持精度,则eps需大于0。默认值为0.00001。
49 49 
50## 返回值说明50## 返回值说明
51 51 
52-- **out** (`Tensor`):数据类型和shape与`input`相同,支持$ND$,支持非连续的Tensor。52+- **out** (`Tensor`):数据类型和shape与`input`相同,支持$ND$,支持非连续的Tensor。
53- - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。53+ - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。
54- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。54+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。
55 55 
56-- **meanOut** (`Tensor`):数据类型与`input`相同,shape为\(N, group\),其中N为`input`第0维度值。数据格式支持$ND$,支持非连续的Tensor。56+- **meanOut** (`Tensor`):数据类型与`input`相同,shape为\(N, group\),其中N为`input`第0维度值。数据格式支持$ND$,支持非连续的Tensor。
57- - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。57+ - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。
58- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。58+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。
59 59 
60-- **rstdOut** (`Tensor`):数据类型与`input`相同,shape为\(N, group\),其中N为`input`第0维度值。数据格式支持$ND$,支持非连续的Tensor。60+- **rstdOut** (`Tensor`):数据类型与`input`相同,shape为\(N, group\),其中N为`input`第0维度值。数据格式支持$ND$,支持非连续的Tensor。
61- - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。61+ - <term>Atlas 推理系列产品</term>:数据类型支持`float16`、`float32`。
62- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。62+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float16`、`float32`、`bfloat16`。
63 63 
64## 约束说明64## 约束说明
65 65 
66-- 该接口支持推理、训练场景下使用。66+- 该接口支持推理、训练场景下使用。
67-- `input`、`weight`、`bias`、`out`、`meanOut`、`rstdOut`数据类型必须在支持的范围之内。67+- `input`、`weight`、`bias`、`out`、`meanOut`、`rstdOut`数据类型必须在支持的范围之内。
68-- `out`、`meanOut`、`rstdOut`的数据类型与`input`相同;`weight`、`bias`与`input`可以不同。68+- `out`、`meanOut`、`rstdOut`的数据类型与`input`相同;`weight`、`bias`与`input`可以不同。
69-- `weight`与`bias`的数据类型必须保持一致,且数据类型的精度不能低于`input`的数据类型。69+- `weight`与`bias`的数据类型必须保持一致,且数据类型的精度不能低于`input`的数据类型。
70-- `weight`与`bias`的维度需为1且元素数量需与输入`input`的第1维度保持相同。70+- `weight`与`bias`的维度需为1且元素数量需与输入`input`的第1维度保持相同。
71-- `input`维度需大于一维且小于等于八维,且`input`第1维度能整除`group`。71+- `input`维度需大于一维且小于等于八维,且`input`第1维度能整除`group`。
72-- `input`任意维都需大于0。72+- `input`任意维都需大于0。
73-- `out`的shape与`input`相同。73+- `out`的shape与`input`相同。
74-- `meanOut`与`rstdOut`的shape为\(N, group\),其中N为`input`第0维度值。74+- `meanOut`与`rstdOut`的shape为\(N, group\),其中N为`input`第0维度值。
75-- `eps`需大于0。75+- `eps`需大于0。
76-- `group`需大于0。76+- `group`需大于0。
77 77 
78## 调用示例78## 调用示例
79 79 
@@ -104,4 +104,3 @@ weight_npu=torch.randn(shape_c,dtype=torch.float16).npu()
104bias_npu=torch.randn(shape_c,dtype=torch.float16).npu()104bias_npu=torch.randn(shape_c,dtype=torch.float16).npu()
105out_npu, mean_npu, rstd_out = torch_npu.npu_group_norm_silu(input_npu, weight_npu, bias_npu, group=num_groups, eps=eps)105out_npu, mean_npu, rstd_out = torch_npu.npu_group_norm_silu(input_npu, weight_npu, bias_npu, group=num_groups, eps=eps)
106```106```
107- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_group_norm_swish.md+8-11
@@ -7,7 +7,6 @@
7|<term>Atlas A3 训练系列产品</term> | √ |7|<term>Atlas A3 训练系列产品</term> | √ |
8|<term>Atlas A2 训练系列产品</term> | √ |8|<term>Atlas A2 训练系列产品</term> | √ |
9 9 
10- 
11## 功能说明10## 功能说明
12 11 
13- API功能:计算输入`input`的组归一化结果`y`,均值`mean`,标准差的倒数`rstd`,以及swish的输出。12- API功能:计算输入`input`的组归一化结果`y`,均值`mean`,标准差的倒数`rstd`,以及swish的输出。
@@ -26,24 +25,23 @@
26 y = \frac{x}{1 + e^{-scale \cdot x}}25 y = \frac{x}{1 + e^{-scale \cdot x}}
27 $$26 $$
28 27
29- 
30> **说明:**<br>28> **说明:**<br>
31-> 需要计算反向梯度场景时,若需要输出结果排除随机性,则需要[设置确定性计算开关](determin_API_list.md)。29+> 需要计算反向梯度场景时,若需要输出结果排除随机性,则需要[设置确定性计算开关](../determin_API_list.md)。
32 30 
33## 函数原型31## 函数原型
34 32 
35-```33+```python
36torch_npu.npu_group_norm_swish(input, num_groups, weight, bias, eps=1e-5, swish_scale=1.0) -> (Tensor, Tensor, Tensor)34torch_npu.npu_group_norm_swish(input, num_groups, weight, bias, eps=1e-5, swish_scale=1.0) -> (Tensor, Tensor, Tensor)
37```35```
38 36 
39## 参数说明37## 参数说明
40 38 
41-- **input**(`Tensor`):必选参数,表示需要进行组归一化的数据,支持2-8D张量,数据类型支持`float16`,`float32`,`bfloat16`。39+- **input**(`Tensor`):必选参数,表示需要进行组归一化的数据,支持2-8D张量,数据类型支持`float16`,`float32`,`bfloat16`。
42-- **num_groups**(`int`):必选参数,表示将`input`的第1维分为`num_groups`组,`input`的第1维必须能被`num_groups`整除。40+- **num_groups**(`int`):必选参数,表示将`input`的第1维分为`num_groups`组,`input`的第1维必须能被`num_groups`整除。
43-- **weight**(`Tensor`):必选参数,表示权重,支持1D张量,并且第0维大小与`input`的第1维相同;数据类型支持`float16`,`float32`,`bfloat16`,并且需要与`input`一致。41+- **weight**(`Tensor`):必选参数,表示权重,支持1D张量,并且第0维大小与`input`的第1维相同;数据类型支持`float16`,`float32`,`bfloat16`,并且需要与`input`一致。
44-- **bias**(`Tensor`):必选参数,表示偏置,支持1D张量,并且第0维大小与`input`的第1维相同;数据类型支持`float16`,`float32`,`bfloat16`,并且需要与`input`一致。42+- **bias**(`Tensor`):必选参数,表示偏置,支持1D张量,并且第0维大小与`input`的第1维相同;数据类型支持`float16`,`float32`,`bfloat16`,并且需要与`input`一致。
45-- **eps**(`float`):可选参数,计算组归一化时加到分母上的值,以保证数值的稳定性。默认值为1e-5。43+- **eps**(`float`):可选参数,计算组归一化时加到分母上的值,以保证数值的稳定性。默认值为1e-5。
46-- **swish_scale**(`float`):可选参数,用于进行swish计算的值。默认值为1.0。44+- **swish_scale**(`float`):可选参数,用于进行swish计算的值。默认值为1.0。
47 45 
48## 返回值说明46## 返回值说明
49 47 
@@ -71,4 +69,3 @@ eps = 1e-5
71swish_scale = 1.069swish_scale = 1.0
72 out, mean, rstd = torch_npu.npu_group_norm_swish(input, num_groups, weight, bias, eps=eps, swish_scale=swish_scale)70 out, mean, rstd = torch_npu.npu_group_norm_swish(input, num_groups, weight, bias, eps=eps, swish_scale=swish_scale)
73```71```
74- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_group_quant.md+2-3
@@ -17,7 +17,7 @@
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu_group_quant(x, scale, group_index, *, offset=None, dst_dtype=None) -> Tensor21torch_npu.npu_group_quant(x, scale, group_index, *, offset=None, dst_dtype=None) -> Tensor
22```22```
23 23 
@@ -31,13 +31,13 @@ torch_npu.npu_group_quant(x, scale, group_index, *, offset=None, dst_dtype=None)
31- **dst_dtype** (`ScalarType`):可选参数,输入值允许为`int8`或`quint4x2`,默认值为`int8`。31- **dst_dtype** (`ScalarType`):可选参数,输入值允许为`int8`或`quint4x2`,默认值为`int8`。
32 32 
33## 返回值说明33## 返回值说明
34+ 
34`Tensor`35`Tensor`
35 36 
36代表`npu_group_quant`的计算结果,对应公式中的`y`。如果参数`dst_dtype``int8`,输出shape与输入`x`的shape一致。如果参数`dst_dtype``quint4x2`,输出的数据类型是`int32`,shape的第0维大小与输入`x`的第0维大小一致,最后一维是输入`x`的最后一维的1/8。支持空Tensor,支持非连续的Tensor。37代表`npu_group_quant`的计算结果,对应公式中的`y`。如果参数`dst_dtype``int8`,输出shape与输入`x`的shape一致。如果参数`dst_dtype``quint4x2`,输出的数据类型是`int32`,shape的第0维大小与输入`x`的第0维大小一致,最后一维是输入`x`的最后一维的1/8。支持空Tensor,支持非连续的Tensor。
37 38 
38## 约束说明39## 约束说明
39 40 
40- 
41- 输入`group_index`必须是非递减序列,最小值不能小于0,最大值必须与输入`x`的shape的第0维大小相等。41- 输入`group_index`必须是非递减序列,最小值不能小于0,最大值必须与输入`x`的shape的第0维大小相等。
42- 该接口支持图模式。42- 该接口支持图模式。
43 43 
@@ -125,4 +125,3 @@ torch_npu.npu_group_quant(x, scale, group_index, *, offset=None, dst_dtype=None)
125 [ 1, -1, 0, 1],125 [ 1, -1, 0, 1],
126 [ 0, 0, -1, -1]], device='npu:0', dtype=torch.int8)126 [ 0, 0, -1, -1]], device='npu:0', dtype=torch.int8)
127 ```127 ```
128- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_grouped_matmul.md+115-112
@@ -10,196 +10,198 @@
10 10 
11## 功能说明<a name="zh-cn_topic_0000002262888689_section1290611593405"></a>11## 功能说明<a name="zh-cn_topic_0000002262888689_section1290611593405"></a>
12 12 
13-- API功能:`npu_grouped_matmul`是一种对多个矩阵乘法(matmul)操作进行分组计算的高效方法。该API实现了对多个矩阵乘法操作的批量处理,通过将具有相同形状或相似形状的矩阵乘法操作组合在一起,减少内存访问开销和计算资源的浪费,从而提高计算效率。13+- API功能:`npu_grouped_matmul`是一种对多个矩阵乘法(matmul)操作进行分组计算的高效方法。该API实现了对多个矩阵乘法操作的批量处理,通过将具有相同形状或相似形状的矩阵乘法操作组合在一起,减少内存访问开销和计算资源的浪费,从而提高计算效率。
14 14 
15-- 计算公式:15+- 计算公式:
16 16 
17 公式中$@$符号表示矩阵乘法,$\times$符号表示矩阵Hadamard乘积:17 公式中$@$符号表示矩阵乘法,$\times$符号表示矩阵Hadamard乘积:
18 18 
19- - 非量化场景(公式1):19+ - 非量化场景(公式1):
20 20 
21 $y_i = x_i @ weight_i + bias_i$21 $y_i = x_i @ weight_i + bias_i$
22 22 
23- - perchannel、pertensor量化场景(公式2):23+ - perchannel、pertensor量化场景(公式2):
24 24 
25 $y_i = (x_i @ weight_i) \times scale_i + offset_i$25 $y_i = (x_i @ weight_i) \times scale_i + offset_i$
26 26 
27- - `x`为`int8`输入,`bias`为`int32`输入(公式2-1):27+ - `x`为`int8`输入,`bias`为`int32`输入(公式2-1):
28 28 
29 $y_i = (x_i @ weight_i + bias_i) \times scale_i + offset_i$29 $y_i = (x_i @ weight_i + bias_i) \times scale_i + offset_i$
30 30 
31- - `x`为`int8`输入,`bias`为`bfloat16`、`float16`、`float32`输入,无offset(公式2-2):31+ - `x`为`int8`输入,`bias`为`bfloat16`、`float16`、`float32`输入,无offset(公式2-2):
32 32 
33 $y_i = (x_i @ weight_i) \times scale_i + bias_i$33 $y_i = (x_i @ weight_i) \times scale_i + bias_i$
34 34 
35- - pertoken、pertensor+pertensor、pertensor+perchannel量化场景(公式3):35+ - pertoken、pertensor+pertensor、pertensor+perchannel量化场景(公式3):
36 36 
37 $y_i = (x_i @ weight_i + bias_i) \times scale_i \times pertokenscale_i$37 $y_i = (x_i @ weight_i + bias_i) \times scale_i \times pertokenscale_i$
38 38 
39- - `x`为`int8`输入,bias为`int32`输入(公式3-1):39+ - `x`为`int8`输入,bias为`int32`输入(公式3-1):
40 40 
41 $y_i = (x_i @ weight_i + bias_i) \times scale_i \times pertokenscale_i$41 $y_i = (x_i @ weight_i + bias_i) \times scale_i \times pertokenscale_i$
42 42 
43- - `x`为`int8`输入,`bias`为`bfloat16`,`float16`,`float32`输入(公式3-2):43+ - `x`为`int8`输入,`bias`为`bfloat16`,`float16`,`float32`输入(公式3-2):
44 44 
45 $y_i = (x_i @ weight_i) \times scale_i \times pertokenscale_i + bias_i$45 $y_i = (x_i @ weight_i) \times scale_i \times pertokenscale_i + bias_i$
46- - `x`为`int4`输入, `weight`的数据类型为`int4`,数据排布格式为`NZ`的输入(公式3-3):46+ - `x`为`int4`输入, `weight`的数据类型为`int4`,数据排布格式为`NZ`的输入(公式3-3):
47 47 
48 $y_i=x_i@ (weight_i \times scale_i) \times pertokenscale_i$48 $y_i=x_i@ (weight_i \times scale_i) \times pertokenscale_i$
49 49
50- 50+ - 伪量化场景(公式4):
51- - 伪量化场景(公式4):
52 51 
53 $y_i = x_i @ ((weight_i + antiquant\_offset_i) \times antiquant\_scale_i) + bias_i$52 $y_i = x_i @ ((weight_i + antiquant\_offset_i) \times antiquant\_scale_i) + bias_i$
54 53 
55## 函数原型<a name="zh-cn_topic_0000002262888689_section87878612417"></a>54## 函数原型<a name="zh-cn_topic_0000002262888689_section87878612417"></a>
56-```55+ 
56+```python
57npu_grouped_matmul(x, weight, *, bias=None, scale=None, offset=None, antiquant_scale=None, antiquant_offset=None, per_token_scale=None, group_list=None, activation_input=None, activation_quant_scale=None, activation_quant_offset=None, split_item=0, group_type=None, group_list_type=0, act_type=0, output_dtype=None, tuning_config=None) -> List[Tensor]57npu_grouped_matmul(x, weight, *, bias=None, scale=None, offset=None, antiquant_scale=None, antiquant_offset=None, per_token_scale=None, group_list=None, activation_input=None, activation_quant_scale=None, activation_quant_offset=None, split_item=0, group_type=None, group_list_type=0, act_type=0, output_dtype=None, tuning_config=None) -> List[Tensor]
58```58```
59 59 
60## 参数说明<a name="zh-cn_topic_0000002262888689_section135561610204110"></a>60## 参数说明<a name="zh-cn_topic_0000002262888689_section135561610204110"></a>
61 61 
62- **x** (`List[Tensor]`):必选参数。输入矩阵列表,表示矩阵乘法中的左矩阵。62- **x** (`List[Tensor]`):必选参数。输入矩阵列表,表示矩阵乘法中的左矩阵。
63- - 支持的数据类型如下:63+ - 支持的数据类型如下:
64- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`float16`、`float32`、`bfloat16`、`int8`和`int4`。64+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`float16`、`float32`、`bfloat16`、`int8`和`int4`。
65- - <term>Atlas 推理系列产品</term>:`float16`。65+ - <term>Atlas 推理系列产品</term>:`float16`。
66 66 
67- - 列表最大长度为128。67+ - 列表最大长度为128。
68- - 当split\_item=0时,张量支持2至6维输入;其他情况下,张量仅支持2维输入。68+ - 当split\_item=0时,张量支持2至6维输入;其他情况下,张量仅支持2维输入。
69 69 
70- **weight** (`List[Tensor]`):必选参数。权重矩阵列表,表示矩阵乘法中的右矩阵。70- **weight** (`List[Tensor]`):必选参数。权重矩阵列表,表示矩阵乘法中的右矩阵。
71- - 支持的数据类型如下:71+ - 支持的数据类型如下:
72- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:72+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
73- - 当`group_list`输入类型为`List[int]`时,支持`float16`、`float32`、`bfloat16`和`int8`。73+ - 当`group_list`输入类型为`List[int]`时,支持`float16`、`float32`、`bfloat16`和`int8`。
74- - 当`group_list`输入类型为`Tensor`时,支持`float16`、`float32`、`bfloat16`、`int4`和`int8`。74+ - 当`group_list`输入类型为`Tensor`时,支持`float16`、`float32`、`bfloat16`、`int4`和`int8`。
75 75 
76- - <term>Atlas 推理系列产品</term>:`float16`。76+ - <term>Atlas 推理系列产品</term>:`float16`。
77 77 
78- - 列表最大长度为128。78+ - 列表最大长度为128。
79- - 每个张量支持2维或3维输入。79+ - 每个张量支持2维或3维输入。
80 80 
81- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。81- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
82- **bias** (`List[Tensor]`):可选参数。每个分组的矩阵乘法输出的独立偏置项。82- **bias** (`List[Tensor]`):可选参数。每个分组的矩阵乘法输出的独立偏置项。
83- - 支持的数据类型如下:83+ - 支持的数据类型如下:
84- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`float16`、`float32`和`int32`。84+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`float16`、`float32`和`int32`。
85- - <term>Atlas 推理系列产品</term>:`float16`。85+ - <term>Atlas 推理系列产品</term>:`float16`。
86 86 
87- - 列表长度与weight列表长度相同。87+ - 列表长度与weight列表长度相同。
88- - 每个张量仅支持1维输入。88+ - 每个张量仅支持1维输入。
89 89 
90- **scale** (`List[Tensor]`):可选参数。用于缩放原数值以匹配量化后的范围值,代表量化参数中的缩放因子,对应公式(2)、公式(3)。90- **scale** (`List[Tensor]`):可选参数。用于缩放原数值以匹配量化后的范围值,代表量化参数中的缩放因子,对应公式(2)、公式(3)。
91- - 支持的数据类型如下:91+ - 支持的数据类型如下:
92- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:92+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
93- - 当`group_list`输入类型为`List[int]`时,支持`int64`。93+ - 当`group_list`输入类型为`List[int]`时,支持`int64`。
94- - 当`group_list`输入类型为`Tensor`时,支持`float32`、`bfloat16`和`int64`。94+ - 当`group_list`输入类型为`Tensor`时,支持`float32`、`bfloat16`和`int64`。
95 95 
96- - <term>Atlas 推理系列产品</term>:仅支持传入`None`。96+ - <term>Atlas 推理系列产品</term>:仅支持传入`None`。
97 97 
98- - 列表长度与weight列表长度相同。98+ - 列表长度与weight列表长度相同。
99- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:每个张量仅支持1维输入。99+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:每个张量仅支持1维输入。
100 100 
101- **offset** (`List[Tensor]`):可选参数。用于调整量化后的数值偏移量,从而更准确地表示原始浮点数值,对应公式(2)。当前仅支持传入`None`101- **offset** (`List[Tensor]`):可选参数。用于调整量化后的数值偏移量,从而更准确地表示原始浮点数值,对应公式(2)。当前仅支持传入`None`
102- **antiquant_scale** (`List[Tensor]`):可选参数。用于缩放原数值以匹配伪量化后的范围值,代表伪量化参数中的缩放因子,对应公式(4)。102- **antiquant_scale** (`List[Tensor]`):可选参数。用于缩放原数值以匹配伪量化后的范围值,代表伪量化参数中的缩放因子,对应公式(4)。
103- - 支持的数据类型如下:103+ - 支持的数据类型如下:
104- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`float16`、`bfloat16`。104+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`float16`、`bfloat16`。
105- - <term>Atlas 推理系列产品</term>:仅支持传入`None`。105+ - <term>Atlas 推理系列产品</term>:仅支持传入`None`。
106 106 
107- - 列表长度与weight列表长度相同。107+ - 列表长度与weight列表长度相同。
108- - 每个张量支持输入维度如下(其中$g$为matmul组数,$G$为pergroup数,$G_i$为第i个tensor的pergroup数):108+ - 每个张量支持输入维度如下(其中$g$为matmul组数,$G$为pergroup数,$G_i$为第i个tensor的pergroup数):
109- - 伪量化perchannel场景,`weight`为单tensor时,shape限制为$[g, n]$;`weight`为多tensor时,shape限制为$[n_i]$。109+ - 伪量化perchannel场景,`weight`为单tensor时,shape限制为$[g, n]$;`weight`为多tensor时,shape限制为$[n_i]$。
110- - 伪量化pergroup场景,weight为单tensor时,shape限制为$[g, G, n]$; weight为多tensor时,shape限制为$[G_i, n_i]$。110+ - 伪量化pergroup场景,weight为单tensor时,shape限制为$[g, G, n]$; weight为多tensor时,shape限制为$[G_i, n_i]$。
111 111 
112- **antiquant_offset** (`List[Tensor]`):可选参数。用于调整伪量化后的数值偏移量,从而更准确地表示原始浮点数值,对应公式(4)。112- **antiquant_offset** (`List[Tensor]`):可选参数。用于调整伪量化后的数值偏移量,从而更准确地表示原始浮点数值,对应公式(4)。
113- - 支持的数据类型如下:113+ - 支持的数据类型如下:
114- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`float16`、`bfloat16`。114+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`float16`、`bfloat16`。
115- - <term>Atlas 推理系列产品</term>:仅支持传入`None`。115+ - <term>Atlas 推理系列产品</term>:仅支持传入`None`。
116 116 
117- - 列表长度与`weight`列表长度相同。117+ - 列表长度与`weight`列表长度相同。
118- - 每个张量输入维度和`antiquant_scale`输入维度一致。118+ - 每个张量输入维度和`antiquant_scale`输入维度一致。
119 119 
120- **per_token_scale** (`List[Tensor]`):可选参数。用于缩放原数值以匹配量化后的范围值,代表pertoken量化参数中由`x`量化引入的缩放因子,对应公式(3)和公式(5)。120- **per_token_scale** (`List[Tensor]`):可选参数。用于缩放原数值以匹配量化后的范围值,代表pertoken量化参数中由`x`量化引入的缩放因子,对应公式(3)和公式(5)。
121- - `group_list`输入类型为`List[int]`时,当前只支持传入`None`。121+ - `group_list`输入类型为`List[int]`时,当前只支持传入`None`。
122- - `group_list`输入类型为`Tensor`时:122+ - `group_list`输入类型为`Tensor`时:
123- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`float32`。123+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`float32`。
124- - 列表长度与`x`列表长度相同。124+ - 列表长度与`x`列表长度相同。
125- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:每个张量仅支持1维输入。125+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:每个张量仅支持1维输入。
126 126 
127- **group_list** (`List[int]`/`Tensor`):可选参数。用于指定分组的索引,表示x的第0维矩阵乘法的索引情况。数据类型支持`int64`。127- **group_list** (`List[int]`/`Tensor`):可选参数。用于指定分组的索引,表示x的第0维矩阵乘法的索引情况。数据类型支持`int64`。
128- - <term>Atlas 推理系列产品</term>:仅支持<code>**Tensor**</code>类型。仅支持1维输入,长度与<code>weight</code>列表长度相同。128+ - <term>Atlas 推理系列产品</term>:仅支持<code>**Tensor**</code>类型。仅支持1维输入,长度与<code>weight</code>列表长度相同。
129- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持<code>**List[int]**</code>或<code>**Tensor**</code>类型。若为<code>**Tensor**</code>类型,仅支持1维输入,长度与<code>weight</code>列表长度相同。129+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持<code>**List[int]**</code>或<code>**Tensor**</code>类型。若为<code>**Tensor**</code>类型,仅支持1维输入,长度与<code>weight</code>列表长度相同。
130- - 配置值要求如下:130+ - 配置值要求如下:
131- - `group_list`输入类型为`List[int]`时,配置值必须为非负递增数列,且长度不能为1。131+ - `group_list`输入类型为`List[int]`时,配置值必须为非负递增数列,且长度不能为1。
132- - `group_list`输入类型为`Tensor`时:132+ - `group_list`输入类型为`Tensor`时:
133- - 当`group_list_type`为0时,`group_list`必须为非负、单调非递减数列。133+ - 当`group_list_type`为0时,`group_list`必须为非负、单调非递减数列。
134- - 当`group_list_type`为1时,`group_list`必须为非负数列,且长度不能为1。134+ - 当`group_list_type`为1时,`group_list`必须为非负数列,且长度不能为1。
135- - 当`group_list_type`为2时,`group_list` shape为$[E, 2]$,E表示Group大小,数据排布为$[[groupIdx0, groupSize0], [groupIdx1, groupSize1]...]$,其中groupSize为分组轴上每组大小,必须为非负数。135+ - 当`group_list_type`为2时,`group_list` shape为$[E, 2]$,E表示Group大小,数据排布为$[[groupIdx0, groupSize0], [groupIdx1, groupSize1]...]$,其中groupSize为分组轴上每组大小,必须为非负数。
136 136 
137- **activation_input** (`List[Tensor]`):可选参数。代表激活函数的反向输入,当前仅支持传入`None`。137- **activation_input** (`List[Tensor]`):可选参数。代表激活函数的反向输入,当前仅支持传入`None`。
138- **activation_quant_scale** (`List[Tensor]`):可选参数。预留参数,当前只支持传入`None`138- **activation_quant_scale** (`List[Tensor]`):可选参数。预留参数,当前只支持传入`None`
139- **activation_quant_offset** (`List[Tensor]`):可选参数。预留参数,当前只支持传入`None`139- **activation_quant_offset** (`List[Tensor]`):可选参数。预留参数,当前只支持传入`None`
140- **split_item** (`int`):可选参数。用于指定切分模式。数据类型支持`int32`。140- **split_item** (`int`):可选参数。用于指定切分模式。数据类型支持`int32`。
141- - 0、1:输出为多个张量,数量与`weight`相同。141+ - 0、1:输出为多个张量,数量与`weight`相同。
142- - 2、3:输出为单个张量。142+ - 2、3:输出为单个张量。
143 143 
144- **group_type** (`int`):可选参数。代表需要分组的轴。数据类型支持`int32`。144- **group_type** (`int`):可选参数。代表需要分组的轴。数据类型支持`int32`。
145- - `group_list`输入类型为`List[int]`时仅支持传入`None`。145+ - `group_list`输入类型为`List[int]`时仅支持传入`None`。
146 146 
147- - `group_list`输入类型为`Tensor`时,若矩阵乘为$C[m,n]=A[m,k]*B[k,n]$,`group_type`支持的枚举值为:-1代表不分组;0代表m轴分组;2代表k轴分组。147+ - `group_list`输入类型为`Tensor`时,若矩阵乘为$C[m,n]=A[m,k]*B[k,n]$,`group_type`支持的枚举值为:-1代表不分组;0代表m轴分组;2代表k轴分组。
148- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前支持取-1、0、2。148+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前支持取-1、0、2。
149- - <term>Atlas 推理系列产品</term>:当前只支持取0。149+ - <term>Atlas 推理系列产品</term>:当前只支持取0。
150 150 
151- **group_list_type** (`int`):可选参数。代表`group_list`的表达形式。数据类型支持`int32`151- **group_list_type** (`int`):可选参数。代表`group_list`的表达形式。数据类型支持`int32`
152- - `group_list`输入类型为`List[int]`时仅支持传入`None`。152+ - `group_list`输入类型为`List[int]`时仅支持传入`None`。
153 153 
154- - `group_list`输入类型为`Tensor`时可取值0、1或2:154+ - `group_list`输入类型为`Tensor`时可取值0、1或2:
155- - 0:默认值,`group_list`中数值为分组轴大小的cumsum结果(累积和)。155+ - 0:默认值,`group_list`中数值为分组轴大小的cumsum结果(累积和)。
156- - 1:`group_list`中数值为分组轴上每组大小。156+ - 1:`group_list`中数值为分组轴上每组大小。
157- - 2:`group_list` shape为$[E, 2]$,E表示Group大小,数据排布为$[[groupIdx0, groupSize0], [groupIdx1, groupSize1]...]$,其中groupSize为分组轴上每组大小。157+ - 2:`group_list` shape为$[E, 2]$,E表示Group大小,数据排布为$[[groupIdx0, groupSize0], [groupIdx1, groupSize1]...]$,其中groupSize为分组轴上每组大小。
158- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅当`x`和`weight`参数输入类型为`INT8`,并且`group_type`取0(m轴分组)时,支持取2。158+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅当`x`和`weight`参数输入类型为`INT8`,并且`group_type`取0(m轴分组)时,支持取2。
159- - <term>Atlas 推理系列产品</term>:不支持取2。159+ - <term>Atlas 推理系列产品</term>:不支持取2。
160 160
161- **act_type** (`int`):可选参数。代表激活函数类型。数据类型支持`int32`。161- **act_type** (`int`):可选参数。代表激活函数类型。数据类型支持`int32`。
162- - `group_list`输入类型为`List[int]`时仅支持传入`None`。162+ - `group_list`输入类型为`List[int]`时仅支持传入`None`。
163 163 
164- - `group_list`输入类型为`Tensor`时,支持的枚举值包括:0代表不激活;1代表`RELU`激活;2代表`GELU_TANH`激活;3代表暂不支持;4代表`FAST_GELU`激活;5代表`SILU`激活。164+ - `group_list`输入类型为`Tensor`时,支持的枚举值包括:0代表不激活;1代表`RELU`激活;2代表`GELU_TANH`激活;3代表暂不支持;4代表`FAST_GELU`激活;5代表`SILU`激活。
165- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为0-5。165+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为0-5。
166- - <term>Atlas 推理系列产品</term>:当前只支持传入0。166+ - <term>Atlas 推理系列产品</term>:当前只支持传入0。
167 167 
168- **output_dtype** (`torch.dtype`):可选参数。输出数据类型。支持的配置包括:168- **output_dtype** (`torch.dtype`):可选参数。输出数据类型。支持的配置包括:
169- - `None`:默认值,表示输出数据类型与输入`x`的数据类型相同。169+ - `None`:默认值,表示输出数据类型与输入`x`的数据类型相同。
170- - 与输出`y`数据类型一致的类型,具体参考[约束说明](#zh-cn_topic_0000002262888689_section618392112366)。170+ - 与输出`y`数据类型一致的类型,具体参考[约束说明](#zh-cn_topic_0000002262888689_section618392112366)。
171 171 
172- **tuning_config** (`List[int]`):可选参数,数组中的第一个元素表示各个专家处理的token数的预期值,算子tiling时会按照数组中的第一个元素进行最优tiling,性能更优(使用场景参见[约束说明](#zh-cn_topic_0000002262888689_section618392112366));从第二个元素开始预留,用户无须填写,未来会进行扩展。如不使用该参数不传即可。172- **tuning_config** (`List[int]`):可选参数,数组中的第一个元素表示各个专家处理的token数的预期值,算子tiling时会按照数组中的第一个元素进行最优tiling,性能更优(使用场景参见[约束说明](#zh-cn_topic_0000002262888689_section618392112366));从第二个元素开始预留,用户无须填写,未来会进行扩展。如不使用该参数不传即可。
173- - <term>Atlas 推理系列产品</term>:当前暂不支持该参数。173+ - <term>Atlas 推理系列产品</term>:当前暂不支持该参数。
174 174 
175## 返回值说明<a name="zh-cn_topic_0000002262888689_section1558311519405"></a>175## 返回值说明<a name="zh-cn_topic_0000002262888689_section1558311519405"></a>
176 176 
177`List[Tensor]`177`List[Tensor]`
178 178 
179-- 当`split_item`为0或1时,返回的张量数量与`weight`相同。179+- 当`split_item`为0或1时,返回的张量数量与`weight`相同。
180-- 当`split_item`为2或3时,返回的张量数量为1。180+- 当`split_item`为2或3时,返回的张量数量为1。
181 181 
182## 约束说明<a name="zh-cn_topic_0000002262888689_section618392112366"></a>182## 约束说明<a name="zh-cn_topic_0000002262888689_section618392112366"></a>
183 183 
184-- 该接口支持推理场景下使用。184+- 该接口支持推理场景下使用。
185-- 该接口支持图模式。185+- 该接口支持图模式。
186-- <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:内轴限制InnerLimit为65536。186+- <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:内轴限制InnerLimit为65536。
187-- `x`和`weight`中每一组tensor的最后一维大小都应小于InnerLimit。x<sub>i</sub>的最后一维指当x不转置时<code>x<sub>i</sub></code>的K轴或当`x`转置时<code>x<sub>i</sub></code>的$M$轴。<code>weight<sub>i</sub></code>的最后一维指当`weight`不转置时<code>weight<sub>i</sub></code>的$N$轴或当`weight`转置时<code>weight<sub>i</sub></code>的$K$轴。187+- `x`和`weight`中每一组tensor的最后一维大小都应小于InnerLimit。x<sub>i</sub>的最后一维指当x不转置时<code>x<sub>i</sub></code>的K轴或当`x`转置时<code>x<sub>i</sub></code>的$M$轴。<code>weight<sub>i</sub></code>的最后一维指当`weight`不转置时<code>weight<sub>i</sub></code>的$N$轴或当`weight`转置时<code>weight<sub>i</sub></code>的$K$轴。
188 188 
189-- `tuning_config`使用场景限制:189+- `tuning_config`使用场景限制:
190 190 
191 仅在量化场景(输入`int8`,输出为`int32`/`bfloat16`/`float16`/`int8`,数据类型如下表),且为单tensor单专家的场景下使用。191 仅在量化场景(输入`int8`,输出为`int32`/`bfloat16`/`float16`/`int8`,数据类型如下表),且为单tensor单专家的场景下使用。
192- |x| weight|output_dtype|y|192+ 
193+ |x| weight|output_dtype|y|
193 |---------|--------|--------|--------|194 |---------|--------|--------|--------|
194 |`int8`|`int8`|`int8`|`int8`|195 |`int8`|`int8`|`int8`|`int8`|
195 |`int8`|`int8`|`bfloat16`|`bfloat16`|196 |`int8`|`int8`|`bfloat16`|`bfloat16`|
196 |`int8`|`int8`|`float16`|`float16`|197 |`int8`|`int8`|`float16`|`float16`|
197 |`int8`|`int8`|`int32`|`int32`|198 |`int8`|`int8`|`int32`|`int32`|
198 199 
199-- 各场景输入与输出数据类型使用约束:200+- 各场景输入与输出数据类型使用约束:
200- - **`group_list`输入类型为`List[int]`时**,<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>数据类型使用约束。201+ - **`group_list`输入类型为`List[int]`时**,<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>数据类型使用约束。
201 202 
202 **表 1** 数据类型约束203 **表 1** 数据类型约束
204+ 
203 |场景|x|weight|bias|scale|antiquant_scale|antiquant_offset|output_dtype|y|205 |场景|x|weight|bias|scale|antiquant_scale|antiquant_offset|output_dtype|y|
204 |---------|--------|--------|--------|--------|--------|--------|--------|--------|206 |---------|--------|--------|--------|--------|--------|--------|--------|--------|
205 |非量化|`float16`|`float16`|`float16`|无需赋值|无需赋值|无需赋值|`float16`|`float16`|207 |非量化|`float16`|`float16`|`float16`|无需赋值|无需赋值|无需赋值|`float16`|`float16`|
@@ -209,10 +211,11 @@ npu_grouped_matmul(x, weight, *, bias=None, scale=None, offset=None, antiquant_s
209 |伪量化|`float16`|`int8`|`float16`|无需赋值|`float16`|`float16`|`float16`|`float16`|211 |伪量化|`float16`|`int8`|`float16`|无需赋值|`float16`|`float16`|`float16`|`float16`|
210 |伪量化|`bfloat16`|`int8`|`float32`|无需赋值|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|212 |伪量化|`bfloat16`|`int8`|`float32`|无需赋值|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|
211 213 
212- - **`group_list`输入类型为`Tensor`时**,数据类型使用约束。214+ - **`group_list`输入类型为`Tensor`时**,数据类型使用约束。
213- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:215+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
214 216 
215 **表 2** 数据类型约束217 **表 2** 数据类型约束
218+ 
216 |场景|x|weight|bias|scale|antiquant_scale|antiquant_offset|per_token_scale|output_dtype|y|219 |场景|x|weight|bias|scale|antiquant_scale|antiquant_offset|per_token_scale|output_dtype|y|
217 |---------|--------|--------|--------|--------|--------|--------|--------|--------|--------|220 |---------|--------|--------|--------|--------|--------|--------|--------|--------|--------|
218 |非量化|`float16`|`float16`|`float16`|无需赋值|无需赋值|无需赋值|无需赋值|None/`float16`|`float16`|221 |非量化|`float16`|`float16`|`float16`|无需赋值|无需赋值|无需赋值|无需赋值|None/`float16`|`float16`|
@@ -229,19 +232,20 @@ npu_grouped_matmul(x, weight, *, bias=None, scale=None, offset=None, antiquant_s
229 |伪量化|`bfloat16`|`int8`/`int4`|`float32`|无需赋值|`bfloat16`|`bfloat16`|无需赋值|None/`bfloat16`|`bfloat16`|232 |伪量化|`bfloat16`|`int8`/`int4`|`float32`|无需赋值|`bfloat16`|`bfloat16`|无需赋值|None/`bfloat16`|`bfloat16`|
230 233 
231 > [!NOTE] 234 > [!NOTE]
232- > - 伪量化场景,若`weight`的类型为`int8`,仅支持perchannel模式;若`weight`的类型为`int4`,支持perchannel和pergroup两种模式。若为pergroup,pergroup数$G$或$G_i$必须要能整除对应的$k_i$。若`weight`为多tensor,定义pergroup长度$s_i= k_i/G_i$,要求所有$s_i(i=1,2,...g)$都相等。235+ > - 伪量化场景,若`weight`的类型为`int8`,仅支持perchannel模式;若`weight`的类型为`int4`,支持perchannel和pergroup两种模式。若为pergroup,pergroup数$G$或$G_i$必须要能整除对应的$k_i$。若`weight`为多tensor,定义pergroup长度$s_i= k_i/G_i$,要求所有$s_i(i=1,2,...g)$都相等。
233- > - 伪量化场景,若`weight`的类型为`int4`,则`weight`中每一组tensor的最后一维大小都应是偶数。<code>weight<sub>i</sub></code>的最后一维指`weight`不转置时<code>weight<sub>i</sub></code>的N轴或当weight转置时weight<sub>i</sub>的$K$轴。并且在pergroup场景下,当`weight`转置时,要求pergroup长度$s_i$是偶数。tensor转置:指若tensor shape为$[M,K]$时,则stride为$[1,M]$,数据排布为$[K,M]$的场景,即非连续tensor。236+ > - 伪量化场景,若`weight`的类型为`int4`,则`weight`中每一组tensor的最后一维大小都应是偶数。<code>weight<sub>i</sub></code>的最后一维指`weight`不转置时<code>weight<sub>i</sub></code>的N轴或当weight转置时weight<sub>i</sub>的$K$轴。并且在pergroup场景下,当`weight`转置时,要求pergroup长度$s_i$是偶数。tensor转置:指若tensor shape为$[M,K]$时,则stride为$[1,M]$,数据排布为$[K,M]$的场景,即非连续tensor。
234- > - 当前PyTorch不支持`int4`类型数据,需要使用时可以通过[torch\_npu.npu\_quantize](torch_npu-npu_quantize.md)接口使用`int32`数据表示`int4`。237+ > - 当前PyTorch不支持`int4`类型数据,需要使用时可以通过[torch\_npu.npu\_quantize](torch_npu-npu_quantize.md)接口使用`int32`数据表示`int4`。
235 238 
236- - <term>Atlas 推理系列产品</term>:239+ - <term>Atlas 推理系列产品</term>:
237 240 
238 **表 3** 数据类型约束241 **表 3** 数据类型约束
242+ 
239 |x|weight|bias|scale|antiquant_scale|antiquant_offset|per_token_scale|output_dtype|y|243 |x|weight|bias|scale|antiquant_scale|antiquant_offset|per_token_scale|output_dtype|y|
240 |--------|--------|--------|--------|--------|--------|--------|--------|--------|244 |--------|--------|--------|--------|--------|--------|--------|--------|--------|
241 |`float16`|`float16`|`float16`|无需赋值|无需赋值|无需赋值|`float32`|`float16`|`float16`|245 |`float16`|`float16`|`float16`|无需赋值|无需赋值|无需赋值|`float32`|`float16`|`float16`|
242 246
243-- 根据输入`x`、输入`weight`与输出`y`的Tensor数量不同,支持以下几种场景。场景中的“单”表示单个张量,“多”表示多个张量。场景顺序为`x`、`weight`、`y`,例如“单多单”表示`x`为单张量,`weight`为多张量,`y`为单张量。247+- 根据输入`x`、输入`weight`与输出`y`的Tensor数量不同,支持以下几种场景。场景中的“单”表示单个张量,“多”表示多个张量。场景顺序为`x`、`weight`、`y`,例如“单多单”表示`x`为单张量,`weight`为多张量,`y`为单张量。
244- - **`group_list`输入类型为`List[int]`时**,<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>各场景的限制。248+ - **`group_list`输入类型为`List[int]`时**,<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>各场景的限制。
245 249 
246 |支持场景|场景说明|场景限制|250 |支持场景|场景说明|场景限制|
247 |--------|--------|--------|251 |--------|--------|--------|
@@ -250,12 +254,12 @@ npu_grouped_matmul(x, weight, *, bias=None, scale=None, offset=None, antiquant_s
250 |单多多|`x`为单张量,`weight`为多张量,`y`为多张量。|1.仅支持`split_item`为0或1。<br>2.必须传`group_list`,`group_list`的差值需与`y`中tensor的第一维一一对应。<br>3.`x`、`weight`、`y`中tensor需为2维。|254 |单多多|`x`为单张量,`weight`为多张量,`y`为多张量。|1.仅支持`split_item`为0或1。<br>2.必须传`group_list`,`group_list`的差值需与`y`中tensor的第一维一一对应。<br>3.`x`、`weight`、`y`中tensor需为2维。|
251 |多多单|`x`和`weight`为多张量,`y`为单张量。每组矩阵乘法的结果连续存放在同一个张量中。|1.仅支持`split_item`为2或3。<br>2.`x`、`weight`、`y`中tensor需为2维。<br>3.`weight`中每个tensor的N轴必须相等。<br>4.若传入`group_list`,`group_list`的差值需与`x`中tensor的第一维一一对应。|255 |多多单|`x`和`weight`为多张量,`y`为单张量。每组矩阵乘法的结果连续存放在同一个张量中。|1.仅支持`split_item`为2或3。<br>2.`x`、`weight`、`y`中tensor需为2维。<br>3.`weight`中每个tensor的N轴必须相等。<br>4.若传入`group_list`,`group_list`的差值需与`x`中tensor的第一维一一对应。|
252 256
253- - **`group_list`输入类型为`Tensor`时**,各场景的限制。257+ - **`group_list`输入类型为`Tensor`时**,各场景的限制。
254- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:258+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
255 259 
256 > [!NOTE] 260 > [!NOTE]
257- > - 量化、伪量化仅支持`group_type`为-1和0场景。261+ > - 量化、伪量化仅支持`group_type`为-1和0场景。
258- > - 仅pertoken量化场景支持激活函数计算。262+ > - 仅pertoken量化场景支持激活函数计算。
259 263 
260 |group_type|支持场景|场景说明|场景限制|264 |group_type|支持场景|场景说明|场景限制|
261 |--------|--------|--------|--------|265 |--------|--------|--------|--------|
@@ -264,19 +268,19 @@ npu_grouped_matmul(x, weight, *, bias=None, scale=None, offset=None, antiquant_s
264 |0|单多单|`x`为单张量,`weight`为多张量,`y`为单张量。|1.仅支持`split_item`为2或3。<br>2.必须传`group_list`,且当`group_list_type`为0时,最后一个值与`x`中tensor的第一维相等,当`group_list_type`为1时,数值的总和与`x`中tensor的第一维相等且长度最大为128,当`group_list_type`为2时,第二列数值的总和与`x`中tensor的第一维相等且长度最大为128。<br>3.`x`、`weight`、`y`中tensor需为2维。<br>4.`weight`中每个tensor的N轴必须相等。<br>5.支持`weight`转置,但`weight`中每个tensor是否转置需保持统一。<br>6.`x`不支持转置。|268 |0|单多单|`x`为单张量,`weight`为多张量,`y`为单张量。|1.仅支持`split_item`为2或3。<br>2.必须传`group_list`,且当`group_list_type`为0时,最后一个值与`x`中tensor的第一维相等,当`group_list_type`为1时,数值的总和与`x`中tensor的第一维相等且长度最大为128,当`group_list_type`为2时,第二列数值的总和与`x`中tensor的第一维相等且长度最大为128。<br>3.`x`、`weight`、`y`中tensor需为2维。<br>4.`weight`中每个tensor的N轴必须相等。<br>5.支持`weight`转置,但`weight`中每个tensor是否转置需保持统一。<br>6.`x`不支持转置。|
265 |0|多多单|`x`和`weight`为多张量,`y`为单张量。每组矩阵乘法的结果连续存放在同一个张量中。|1.仅支持`split_item`为2或3。<br>2.`x`、`weight`、`y`中tensor需为2维。<br>3.`weight`中每个tensor的N轴必须相等。<br>4.若传入`group_list`,当`group_list_type`为0时,`group_list`的差值需与`x`中tensor的第一维一一对应,当`group_list_type`为1时,`group_list`的数值需与`x`中tensor的第一维一一对应且长度最大为128,当`group_list_type`为2时,`group_list`第二列的数值需与`x`中tensor的第一维一一对应且长度最大为128。<br>5.支持`weight`转置,但`weight`中每个tensor是否转置需保持统一。<br>6.`x`不支持转置。|269 |0|多多单|`x`和`weight`为多张量,`y`为单张量。每组矩阵乘法的结果连续存放在同一个张量中。|1.仅支持`split_item`为2或3。<br>2.`x`、`weight`、`y`中tensor需为2维。<br>3.`weight`中每个tensor的N轴必须相等。<br>4.若传入`group_list`,当`group_list_type`为0时,`group_list`的差值需与`x`中tensor的第一维一一对应,当`group_list_type`为1时,`group_list`的数值需与`x`中tensor的第一维一一对应且长度最大为128,当`group_list_type`为2时,`group_list`第二列的数值需与`x`中tensor的第一维一一对应且长度最大为128。<br>5.支持`weight`转置,但`weight`中每个tensor是否转置需保持统一。<br>6.`x`不支持转置。|
266 270 
267- 271+ - <term>Atlas 推理系列产品</term>:
268- - <term>Atlas 推理系列产品</term>
269 272 
270 输入输出只支持`float16`的数据类型,输出`y`的n轴大小需要是16的倍数。273 输入输出只支持`float16`的数据类型,输出`y`的n轴大小需要是16的倍数。
274+ 
271 |group_type|支持场景|场景说明|场景限制|275 |group_type|支持场景|场景说明|场景限制|
272 |--------|--------|--------|--------|276 |--------|--------|--------|--------|
273 |0|单单单|`x`、`weight`与`y`均为单张量。|1.仅支持`split_item`为2或3。<br>2.`weight`中tensor需为3维,`x`、`y`中tensor需为2维。<br>3.必须传`group_list`,且当`group_list_type`为0时,最后一个值与`x`中tensor的第一维相等,当`group_list_type`为1时,数值的总和与`x`中tensor的第一维相等。<br>4.`group_list`第1维最大支持1024,即最多支持1024个group。<br>5.支持`weight`转置,不支持`x`转置。|277 |0|单单单|`x`、`weight`与`y`均为单张量。|1.仅支持`split_item`为2或3。<br>2.`weight`中tensor需为3维,`x`、`y`中tensor需为2维。<br>3.必须传`group_list`,且当`group_list_type`为0时,最后一个值与`x`中tensor的第一维相等,当`group_list_type`为1时,数值的总和与`x`中tensor的第一维相等。<br>4.`group_list`第1维最大支持1024,即最多支持1024个group。<br>5.支持`weight`转置,不支持`x`转置。|
274 278 
275## 调用示例<a name="zh-cn_topic_0000002262888689_section1566973054111"></a>279## 调用示例<a name="zh-cn_topic_0000002262888689_section1566973054111"></a>
276 280 
277-- 单算子模式调用281+- 单算子模式调用
278 282 
279- - 通用调用示例283+ - 通用调用示例
280 284 
281 ```python285 ```python
282 import torch286 import torch
@@ -302,7 +306,7 @@ npu_grouped_matmul(x, weight, *, bias=None, scale=None, offset=None, antiquant_s
302 npu_out = torch_npu.npu_grouped_matmul(x, weight, bias=bias, group_list=group_list, split_item=split_item, group_type=-1)306 npu_out = torch_npu.npu_grouped_matmul(x, weight, bias=bias, group_list=group_list, split_item=split_item, group_type=-1)
303 ```307 ```
304 308 
305- - x为int4输入, weight的数据类型为int4数据排布格式为NZ,调用示例如下:309+ - x为int4输入, weight的数据类型为int4数据排布格式为NZ,调用示例如下:
306 310 
307 ```python311 ```python
308 import numpy as np312 import numpy as np
@@ -336,8 +340,8 @@ npu_grouped_matmul(x, weight, *, bias=None, scale=None, offset=None, antiquant_s
336 split_item=3, output_dtype=torch.float16)340 split_item=3, output_dtype=torch.float16)
337 ```341 ```
338 342 
339-- 图模式调用343+- 图模式调用
340- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:344+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>/<term>Atlas 推理系列产品</term>/<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
341 345 
342 ```python346 ```python
343 import torch347 import torch
@@ -374,4 +378,3 @@ npu_grouped_matmul(x, weight, *, bias=None, scale=None, offset=None, antiquant_s
374 if __name__ == '__main__':378 if __name__ == '__main__':
375 main()379 main()
376 ```380 ```
377- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_grouped_matmul_finalize_routing.md+9-9
@@ -13,7 +13,7 @@ GroupedMatMul和MoeFinalizeRouting的融合算子,GroupedMatMul计算后的输
13 13 
14## 函数原型<a name="zh-cn_topic_0000002259406069_section45077510411"></a>14## 函数原型<a name="zh-cn_topic_0000002259406069_section45077510411"></a>
15 15 
16-```16+```python
17torch_npu.npu_grouped_matmul_finalize_routing(x, w, group_list, *, scale=None, bias=None, offset=None, pertoken_scale=None, shared_input=None, logit=None, row_index=None, dtype=None, shared_input_weight=1.0, shared_input_offset=0, output_bs=0, group_list_type=1) -> Tensor17torch_npu.npu_grouped_matmul_finalize_routing(x, w, group_list, *, scale=None, bias=None, offset=None, pertoken_scale=None, shared_input=None, logit=None, row_index=None, dtype=None, shared_input_weight=1.0, shared_input_offset=0, output_bs=0, group_list_type=1) -> Tensor
18```18```
19 19 
@@ -21,8 +21,8 @@ torch_npu.npu_grouped_matmul_finalize_routing(x, w, group_list, *, scale=None, b
21 21 
22- **x** (`Tensor`):必选参数。矩阵计算的左矩阵,不支持非连续的Tensor。数据类型支持`int8`,数据格式支持$ND$,维度为\(m, k\)。m取值范围为\[1, 16\*1024\*8\]。22- **x** (`Tensor`):必选参数。矩阵计算的左矩阵,不支持非连续的Tensor。数据类型支持`int8`,数据格式支持$ND$,维度为\(m, k\)。m取值范围为\[1, 16\*1024\*8\]。
23- **w** (`Tensor`):必选参数。矩阵计算的右矩阵,不支持非连续的Tensor。数据类型支持`int8``int4`23- **w** (`Tensor`):必选参数。矩阵计算的右矩阵,不支持非连续的Tensor。数据类型支持`int8``int4`
24- - A8W8量化场景下,数据格式支持$NZ$,维度为\(e, n1, k1, k0, n0\),其中k0=16、n0=32,`x` shape中的k和`w` shape中的k1需要满足以下关系:ceilDiv\(k, 16\) = k1,e取值范围\[1, 256\],k取值为16整倍数,n取值为32整倍数,且n大于等于256。24+ - A8W8量化场景下,数据格式支持$NZ$,维度为\(e, n1, k1, k0, n0\),其中k0=16、n0=32,`x` shape中的k和`w` shape中的k1需要满足以下关系:ceilDiv\(k, 16\) = k1,e取值范围\[1, 256\],k取值为16整倍数,n取值为32整倍数,且n大于等于256。
25- - A8W4量化场景下数据格式支持$ND$,维度为\(e, k, n\),k支持2048,n只支持7168。25+ - A8W4量化场景下数据格式支持$ND$,维度为\(e, k, n\),k支持2048,n只支持7168。
26 26 
27- **group\_list** (`Tensor`):必选参数。GroupedMatMul的各分组大小。不支持非连续的Tensor。数据类型支持`int64`,数据格式支持$ND$,维度为\(e,\),e与`w`的e一致。`group_list`的值总和要求≤m。27- **group\_list** (`Tensor`):必选参数。GroupedMatMul的各分组大小。不支持非连续的Tensor。数据类型支持`int64`,数据格式支持$ND$,维度为\(e,\),e与`w`的e一致。`group_list`的值总和要求≤m。
28- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。28- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
@@ -47,9 +47,10 @@ torch_npu.npu_grouped_matmul_finalize_routing(x, w, group_list, *, scale=None, b
47 47 
48## 约束说明<a name="zh-cn_topic_0000002259406069_section12345537164214"></a>48## 约束说明<a name="zh-cn_topic_0000002259406069_section12345537164214"></a>
49 49 
50-- 该接口支持推理和训练场景下使用。50+- 该接口支持推理和训练场景下使用。
51-- 该接口支持图模式。51+- 该接口支持图模式。
52-- 输入和输出Tensor支持的数据类型组合如下:52+- 输入和输出Tensor支持的数据类型组合如下:
53+ 
53 |x|w|group_list|scale|bias|offset|pertoken_scale|shared_input|logit|row_index|y|54 |x|w|group_list|scale|bias|offset|pertoken_scale|shared_input|logit|row_index|y|
54 |--------|--------|--------|--------|--------|--------|--------|--------|--------|--------|--------|55 |--------|--------|--------|--------|--------|--------|--------|--------|--------|--------|--------|
55 |`int8`|`int8`|`int64`|`float32`|None|None|`float32`|`bfloat16`|`float32`|`int64`|`float32`|56 |`int8`|`int8`|`int64`|`float32`|None|None|`float32`|`bfloat16`|`float32`|`int64`|`float32`|
@@ -59,7 +60,7 @@ torch_npu.npu_grouped_matmul_finalize_routing(x, w, group_list, *, scale=None, b
59 60 
60## 调用示例<a name="zh-cn_topic_0000002259406069_section14459801435"></a>61## 调用示例<a name="zh-cn_topic_0000002259406069_section14459801435"></a>
61 62 
62-- 单算子模式调用63+- 单算子模式调用
63 64 
64 ```python65 ```python
65 import numpy as np66 import numpy as np
@@ -101,7 +102,7 @@ torch_npu.npu_grouped_matmul_finalize_routing(x, w, group_list, *, scale=None, b
101 shared_input_offset=shared_input_offset, output_bs=output_bs)102 shared_input_offset=shared_input_offset, output_bs=output_bs)
102 ```103 ```
103 104 
104-- 图模式调用:105+- 图模式调用:
105 106 
106 ```python107 ```python
107 import numpy as np108 import numpy as np
@@ -156,4 +157,3 @@ torch_npu.npu_grouped_matmul_finalize_routing(x, w, group_list, *, scale=None, b
156 model = torch.compile(model, backend=npu_backend, dynamic=False)157 model = torch.compile(model, backend=npu_backend, dynamic=False)
157 y = model(x_clone, weightNz, group_list_clone, scale_clone, pertoken_scale_clone, shared_input_clone, logit_clone, row_index_clone, shared_input_offset, output_bs)158 y = model(x_clone, weightNz, group_list_clone, scale_clone, pertoken_scale_clone, shared_input_clone, logit_clone, row_index_clone, shared_input_offset, output_bs)
158 ```159 ```
159- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_grouped_matmul_swiglu_quant_v2.md+12-12
@@ -113,7 +113,7 @@
113 113 
114## 函数原型114## 函数原型
115 115 
116-```116+```python
117torch_npu.npu_grouped_matmul_swiglu_quant_v2(x, weight, weight_scale, x_scale, group_list, *, smooth_scale=None, weight_assist_matrix=None, bias=None, dequant_mode=0, dequant_dtype=0, quant_mode=0, quant_dtype=0, group_list_type=0, tuning_config=None) -> (Tensor, Tensor)117torch_npu.npu_grouped_matmul_swiglu_quant_v2(x, weight, weight_scale, x_scale, group_list, *, smooth_scale=None, weight_assist_matrix=None, bias=None, dequant_mode=0, dequant_dtype=0, quant_mode=0, quant_dtype=0, group_list_type=0, tuning_config=None) -> (Tensor, Tensor)
118```118```
119 119 
@@ -145,11 +145,11 @@ torch_npu.npu_grouped_matmul_swiglu_quant_v2(x, weight, weight_scale, x_scale, g
145 145 
146## 约束说明146## 约束说明
147 147 
148-- 该接口支持推理和训练场景下使用。148+- 该接口支持推理和训练场景下使用。
149-- 该接口支持图模式。149+- 该接口支持图模式。
150-- 确定性计算:该接口默认为确定性实现,即对于相同的输入,多次执行会产生相同的结果,确保计算结果的可重复性。150+- 确定性计算:该接口默认为确定性实现,即对于相同的输入,多次执行会产生相同的结果,确保计算结果的可重复性。
151-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>、<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:151+- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>、<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
152- - 支持A8W8、A8W4、A4W4量化场景,输入和输出Tensor支持的数据类型组合如下:152+ - 支持A8W8、A8W4、A4W4量化场景,输入和输出Tensor支持的数据类型组合如下:
153 153 
154 |量化场景|x|weight|weight\_scale|x\_scale|smooth\_scale|output|output\_scale|154 |量化场景|x|weight|weight\_scale|x\_scale|smooth\_scale|output|output\_scale|
155 |--------|--------|--------|--------|--------|--------|--------|--------|155 |--------|--------|--------|--------|--------|--------|--------|--------|
@@ -157,7 +157,7 @@ torch_npu.npu_grouped_matmul_swiglu_quant_v2(x, weight, weight_scale, x_scale, g
157 |A8W4|`int8`|`int4`、`int32`|`uint64`|`float32`|-|`int8`|`float32`|157 |A8W4|`int8`|`int4`、`int32`|`uint64`|`float32`|-|`int8`|`float32`|
158 |A4W4|`int4`、`int32`|`int4`、`int32`|`float32`|`float32`|`float32`|`int8`|`float32`|158 |A4W4|`int4`、`int32`|`int4`、`int32`|`float32`|`float32`|`float32`|`int8`|`float32`|
159 159 
160- - shape约束如下:160+ - shape约束如下:
161 161 
162 |量化场景|x|weight|weight\_scale|x\_scale|smooth\_scale|output|output\_scale|162 |量化场景|x|weight|weight\_scale|x\_scale|smooth\_scale|output|output\_scale|
163 |--------|--------|--------|--------|--------|--------|--------|--------|163 |--------|--------|--------|--------|--------|--------|--------|--------|
@@ -165,13 +165,13 @@ torch_npu.npu_grouped_matmul_swiglu_quant_v2(x, weight, weight_scale, x_scale, g
165 |A8W4|(M, K)|$ND$格式{(E, K, N)}或NZ格式|per-channel:{(E, N)}; per-group:{(E, K\_group\_num, N)}|(M,)|-|(M, N/2)|(M,)|165 |A8W4|(M, K)|$ND$格式{(E, K, N)}或NZ格式|per-channel:{(E, N)}; per-group:{(E, K\_group\_num, N)}|(M,)|-|(M, N/2)|(M,)|
166 |A4W4|(M, K)|$ND$格式{(E, K, N)}或NZ格式|{(E, N)}|(M,)|(E, N/2)或(E,)|(M, N/2)|(M,)|166 |A4W4|(M, K)|$ND$格式{(E, K, N)}或NZ格式|{(E, N)}|(M,)|(E, N/2)或(E,)|(M, N/2)|(M,)|
167 167 
168- - A8W8场景下,不支持N轴长度超过10240,不支持`x`的尾轴长度大于等于65536。168+ - A8W8场景下,不支持N轴长度超过10240,不支持`x`的尾轴长度大于等于65536。
169- - A8W4场景下,不支持N轴长度超过10240,不支持`x`的尾轴长度大于等于20000。169+ - A8W4场景下,不支持N轴长度超过10240,不支持`x`的尾轴长度大于等于20000。
170- - A4W4场景下,不支持N轴长度超过10240,不支持`x`的尾轴长度大于等于20000。170+ - A4W4场景下,不支持N轴长度超过10240,不支持`x`的尾轴长度大于等于20000。
171 171 
172## 调用示例172## 调用示例
173 173 
174-- 单算子模式调用174+- 单算子模式调用
175 175 
176 ```python176 ```python
177 import numpy as np177 import numpy as np
@@ -197,7 +197,7 @@ torch_npu.npu_grouped_matmul_swiglu_quant_v2(x, weight, weight_scale, x_scale, g
197 output0_npu, output1_npu = torch_npu.npu_grouped_matmul_swiglu_quant_v2(x.npu(), [weight_npu], [weightScale.npu()], xScale.npu(), groupList.npu())197 output0_npu, output1_npu = torch_npu.npu_grouped_matmul_swiglu_quant_v2(x.npu(), [weight_npu], [weightScale.npu()], xScale.npu(), groupList.npu())
198 ```198 ```
199 199 
200-- 图模式调用:200+- 图模式调用:
201 201 
202 ```python202 ```python
203 import numpy as np203 import numpy as np
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_incre_flash_attention.md+3-2
@@ -17,7 +17,7 @@
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu_incre_flash_attention(query, key, value, *, padding_mask=None, pse_shift=None, atten_mask=None, actual_seq_lengths=None, dequant_scale1=None, quant_scale1=None, dequant_scale2=None, quant_scale2=None, quant_offset2=None, antiquant_scale=None, antiquant_offset=None, block_table=None, kv_padding_size=None, num_heads=1, scale_value=1.0, input_layout="BSH", num_key_value_heads=0, block_size=0, inner_precise=1) -> Tensor21torch_npu.npu_incre_flash_attention(query, key, value, *, padding_mask=None, pse_shift=None, atten_mask=None, actual_seq_lengths=None, dequant_scale1=None, quant_scale1=None, dequant_scale2=None, quant_scale2=None, quant_offset2=None, antiquant_scale=None, antiquant_offset=None, block_table=None, kv_padding_size=None, num_heads=1, scale_value=1.0, input_layout="BSH", num_key_value_heads=0, block_size=0, inner_precise=1) -> Tensor
22```22```
23 23 
@@ -72,9 +72,11 @@ torch_npu.npu_incre_flash_attention(query, key, value, *, padding_mask=None, pse
72- **inner_precise** (`int`):可选参数。代表高精度/高性能选择,`0`代表高精度,`1`代表高性能,默认值为`1`(高性能),数据类型支持`int64`。72- **inner_precise** (`int`):可选参数。代表高精度/高性能选择,`0`代表高精度,`1`代表高性能,默认值为`1`(高性能),数据类型支持`int64`。
73 73 
74## 返回值说明74## 返回值说明
75+ 
75`Tensor`76`Tensor`
76 77 
77计算的最终结果,对应公式中的$atten\_out$,`shape`与`query`保持一致。78计算的最终结果,对应公式中的$atten\_out$,`shape`与`query`保持一致。
79+ 
78- 非量化场景下,输出数据类型与`query`的数据类型保持一致。80- 非量化场景下,输出数据类型与`query`的数据类型保持一致。
79- 量化场景下,若传入`quant_scale2`,则输出数据类型为`int8`81- 量化场景下,若传入`quant_scale2`,则输出数据类型为`int8`
80 82 
@@ -197,4 +199,3 @@ torch_npu.npu_incre_flash_attention(query, key, value, *, padding_mask=None, pse
197 [[-0.9595, -0.9609, -0.6602, ..., 0.7959, 1.7920, 0.0783]]],199 [[-0.9595, -0.9609, -0.6602, ..., 0.7959, 1.7920, 0.0783]]],
198 device='npu:0', dtype=torch.float16) torch.Size([2, 1, 5120])200 device='npu:0', dtype=torch.float16) torch.Size([2, 1, 5120])
199 ```201 ```
200- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_interleave_rope.md+12-13
@@ -9,8 +9,8 @@
9 9 
10## 功能说明10## 功能说明
11 11 
12-- API功能:针对单输入`x`进行旋转位置编码。12+- API功能:针对单输入`x`进行旋转位置编码。
13-- 计算公式:13+- 计算公式:
14 14 
15 ![](../../figures/zh-cn_formulaimage_0000002238091144.png)15 ![](../../figures/zh-cn_formulaimage_0000002238091144.png)
16 16 
@@ -20,15 +20,15 @@
20 20 
21## 函数原型21## 函数原型
22 22 
23-```23+```python
24torch_npu.npu_interleave_rope(x, cos, sin) -> Tensor24torch_npu.npu_interleave_rope(x, cos, sin) -> Tensor
25```25```
26 26 
27## 参数说明27## 参数说明
28 28 
29-- **x** (`Tensor`):表示待处理张量。要求为4维张量,shape为\(B, N, S, D\),数据类型支持`bfloat16`、`float16`,数据格式为$ND$,不支持非连续的Tensor。29+- **x** (`Tensor`):表示待处理张量。要求为4维张量,shape为\(B, N, S, D\),数据类型支持`bfloat16`、`float16`,数据格式为$ND$,不支持非连续的Tensor。
30-- **cos** (`Tensor`):表示RoPE旋转位置编码的余弦分量。要求为4维张量,shape为\(B, N, S, D\),S可以为1或与`x`的S相同,数据类型、数据格式与`x`一致,不支持非连续的Tensor。30+- **cos** (`Tensor`):表示RoPE旋转位置编码的余弦分量。要求为4维张量,shape为\(B, N, S, D\),S可以为1或与`x`的S相同,数据类型、数据格式与`x`一致,不支持非连续的Tensor。
31-- **sin** (`Tensor`):表示RoPE旋转位置编码的正弦分量。shape、数据类型、数据格式需要与`cos`保持一致,不支持非连续的Tensor。31+- **sin** (`Tensor`):表示RoPE旋转位置编码的正弦分量。shape、数据类型、数据格式需要与`cos`保持一致,不支持非连续的Tensor。
32 32 
33## 返回值说明33## 返回值说明
34 34 
@@ -38,14 +38,14 @@ torch_npu.npu_interleave_rope(x, cos, sin) -> Tensor
38 38 
39## 约束说明39## 约束说明
40 40 
41-- 该接口支持推理场景下使用。41+- 该接口支持推理场景下使用。
42-- 该接口支持图模式。42+- 该接口支持图模式。
43-- 输入`x`、`cos`、`sin`的D维度均必须等于64。43+- 输入`x`、`cos`、`sin`的D维度均必须等于64。
44-- `cos`、`sin`的N维度必须等于1。44+- `cos`、`sin`的N维度必须等于1。
45 45 
46## 调用示例46## 调用示例
47 47 
48-- 单算子模式调用48+- 单算子模式调用
49 49 
50 ```python50 ```python
51 import torch51 import torch
@@ -63,7 +63,7 @@ torch_npu.npu_interleave_rope(x, cos, sin) -> Tensor
63 q_embed_npu = torch_npu.npu_interleave_rope(x_npu, cos_npu, sin_npu)63 q_embed_npu = torch_npu.npu_interleave_rope(x_npu, cos_npu, sin_npu)
64 ```64 ```
65 65 
66-- 图模式调用66+- 图模式调用
67 67 
68 ```python68 ```python
69 # 入图方式69 # 入图方式
@@ -97,4 +97,3 @@ torch_npu.npu_interleave_rope(x, cos, sin) -> Tensor
97 # 调用InterleaveRope算子97 # 调用InterleaveRope算子
98 q_embed_npu = model(x_npu, cos_npu, sin_npu)98 q_embed_npu = model(x_npu, cos_npu, sin_npu)
99 ```99 ```
100- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_kv_quant_sparse_flash_attention.md+38-36
@@ -1,6 +1,7 @@
1# torch_npu-npu_kv_quant_sparse_flash_attention1# torch_npu-npu_kv_quant_sparse_flash_attention
2 2 
3## 产品支持情况3## 产品支持情况
4+ 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
5| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
6|<term>Atlas A2 推理系列产品</term> | √ |7|<term>Atlas A2 推理系列产品</term> | √ |
@@ -8,9 +9,9 @@
8 9 
9## 功能说明10## 功能说明
10 11 
11-- API功能:`kv_quant_sparse_flash_attention`在`sparse_flash_attention`的基础上支持了[Per-Token-Head-Tile-128量化]输入。随着大模型上下文长度的增加,Sparse Attention的重要性与日俱增,这一技术通过“只计算关键部分”大幅减少计算量,然而会引入大量的离散访存,造成数据搬运时间增加,进而影响整体性能。12+- API功能:`kv_quant_sparse_flash_attention`在`sparse_flash_attention`的基础上支持了[Per-Token-Head-Tile-128量化]输入。随着大模型上下文长度的增加,Sparse Attention的重要性与日俱增,这一技术通过“只计算关键部分”大幅减少计算量,然而会引入大量的离散访存,造成数据搬运时间增加,进而影响整体性能。
12 13 
13-- 计算公式:14+- 计算公式:
14 $$15 $$
15 Attention=\text{softmax}(\frac{Q @ \text{Dequant}({\tilde{K}^{INT8}},{Scale_K})^T}{\sqrt{d_k}})@\text{Dequant}(\tilde{V}^{INT8},{Scale_V}),16 Attention=\text{softmax}(\frac{Q @ \text{Dequant}({\tilde{K}^{INT8}},{Scale_K})^T}{\sqrt{d_k}})@\text{Dequant}(\tilde{V}^{INT8},{Scale_V}),
16 $$17 $$
@@ -19,81 +20,83 @@
19 20 
20## 函数原型21## 函数原型
21 22 
22-```23+```python
23torch_npu.npu_kv_quant_sparse_flash_attention(query, key, value, sparse_indices, scale_value, key_quant_mode, value_quant_mode, *, key_dequant_scale=None, value_dequant_scale=None, block_table=None, actual_seq_lengths_query=None, actual_seq_lengths_kv=None, sparse_block_size=1, layout_query="BSND", layout_kv="BSND", sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1, attention_mode=0, quant_scale_repo_mode=1, tile_size=128, rope_head_dim=64) -> Tensor24torch_npu.npu_kv_quant_sparse_flash_attention(query, key, value, sparse_indices, scale_value, key_quant_mode, value_quant_mode, *, key_dequant_scale=None, value_dequant_scale=None, block_table=None, actual_seq_lengths_query=None, actual_seq_lengths_kv=None, sparse_block_size=1, layout_query="BSND", layout_kv="BSND", sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1, attention_mode=0, quant_scale_repo_mode=1, tile_size=128, rope_head_dim=64) -> Tensor
24```25```
25 26 
26## 参数说明27## 参数说明
27 28 
28> [!NOTE] 29> [!NOTE]
30+>
29> - query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。31> - query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。
30> - Q\_S和S1表示query shape中的S,KV\_S和S2表示key shape中的S,Q\_N表示num\_query\_heads,KV\_N表示num\_key\_value\_heads,Q\_T表示query shape中的T,KV\_T表示key shape中的T。32> - Q\_S和S1表示query shape中的S,KV\_S和S2表示key shape中的S,Q\_N表示num\_query\_heads,KV\_N表示num\_key\_value\_heads,Q\_T表示query shape中的T,KV\_T表示key shape中的T。
31 33 
32-- **query**(`Tensor`):必选参数,表示attention结构的Q输入,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`和`float16`,query由相同dtype的q_nope和q_rope按D维度拼接得到。`layout_query`为BSND时shape为[B,S1,Q\_N,D],当`layout_query`为TND时shape为[Q\_T,Q\_N,D],其中Q\_N支持1/2/4/8/16/32/64/128。34+- **query**(`Tensor`):必选参数,表示attention结构的Q输入,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`和`float16`,query由相同dtype的q_nope和q_rope按D维度拼接得到。`layout_query`为BSND时shape为[B,S1,Q\_N,D],当`layout_query`为TND时shape为[Q\_T,Q\_N,D],其中Q\_N支持1/2/4/8/16/32/64/128。
33 35 
34-- **key**(`Tensor`):必选参数,表示attention结构的K输入,不支持非连续,数据格式支持$ND$,数据类型支持`int8`,`int8`的k_nope、query相同dtype的k_rope和`float32`的量化参数按D维度拼接得到,layout\_kv为PA\_BSND时shape为[block\_num, block\_size, KV\_N, D],其中block\_num为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的整数倍,最大支持到1024。`layout_kv`为BSND时shape为[B, S2, KV\_N, D],`layout_kv`为TND时shape为[KV\_T, KV\_N, D],其中KV\_N只支持1。36+- **key**(`Tensor`):必选参数,表示attention结构的K输入,不支持非连续,数据格式支持$ND$,数据类型支持`int8`,`int8`的k_nope、query相同dtype的k_rope和`float32`的量化参数按D维度拼接得到,layout\_kv为PA\_BSND时shape为[block\_num, block\_size, KV\_N, D],其中block\_num为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的整数倍,最大支持到1024。`layout_kv`为BSND时shape为[B, S2, KV\_N, D],`layout_kv`为TND时shape为[KV\_T, KV\_N, D],其中KV\_N只支持1。
35 37 
36-- **value**(`Tensor`):必选参数,表示attention结构的V输入,不支持非连续,数据格式支持$ND$,数据类型支持`int8`。value的N仅支持1。38+- **value**(`Tensor`):必选参数,表示attention结构的V输入,不支持非连续,数据格式支持$ND$,数据类型支持`int8`。value的N仅支持1。
37 39 
38-- **sparse\_indices**(`Tensor`):必选参数,代表离散取kvCache的索引,不支持非连续,数据格式支持$ND$,数据类型支持`int32`,当`layout_query`为BSND时,shape需要传入[B, Q\_S, KV\_N, sparse\_size],当`layout_query`为TND时,shape需要传入[Q\_T, KV\_N, sparse\_size],其中sparse\_size为一次离散选取的block数,需要保证每行有效值均在前半部分,无效值均在后半部分,且需要满足sparse\_size大于0。40+- **sparse\_indices**(`Tensor`):必选参数,代表离散取kvCache的索引,不支持非连续,数据格式支持$ND$,数据类型支持`int32`,当`layout_query`为BSND时,shape需要传入[B, Q\_S, KV\_N, sparse\_size],当`layout_query`为TND时,shape需要传入[Q\_T, KV\_N, sparse\_size],其中sparse\_size为一次离散选取的block数,需要保证每行有效值均在前半部分,无效值均在后半部分,且需要满足sparse\_size大于0。
39 41 
40-- **scale\_value**(`double):必选参数,公式中$d_k$开根号的倒数,代表缩放系数,作为query和key矩阵乘后Muls的scalar值,数据类型支持`double`。42+- **scale\_value**(`double):必选参数,公式中$d_k$开根号的倒数,代表缩放系数,作为query和key矩阵乘后Muls的scalar值,数据类型支持`double`。
41 43 
42-- **key\_quant\_mode**(`int`):必选参数,代表key的量化模式,数据类型支持`int64`,仅支持传入2,代表per_tile量化模式。44+- **key\_quant\_mode**(`int`):必选参数,代表key的量化模式,数据类型支持`int64`,仅支持传入2,代表per_tile量化模式。
43 45 
44-- **value\_quant\_mode**(`int`):必选参数,代表value的量化模式,数据类型支持`int64`,仅支持传入2,代表per_tile量化模式。46+- **value\_quant\_mode**(`int`):必选参数,代表value的量化模式,数据类型支持`int64`,仅支持传入2,代表per_tile量化模式。
45 47 
46- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。48- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
47 49 
48-- **key\_dequant\_scale**(`Tensor`):可选参数,预留参数,仅支持默认值。50+- **key\_dequant\_scale**(`Tensor`):可选参数,预留参数,仅支持默认值。
49 51 
50-- **value\_dequant\_scale**(`Tensor`):可选参数,预留参数,仅支持默认值。52+- **value\_dequant\_scale**(`Tensor`):可选参数,预留参数,仅支持默认值。
51 53 
52-- **block\_table**(`Tensor`):可选参数,表示PageAttention中kvCache存储使用的block映射表。数据格式支持$ND$,数据类型支持`int32`,shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的s2对应的block数量,即s2\_max / block\_size向上取整。54+- **block\_table**(`Tensor`):可选参数,表示PageAttention中kvCache存储使用的block映射表。数据格式支持$ND$,数据类型支持`int32`,shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的s2对应的block数量,即s2\_max / block\_size向上取整。
53 55 
54-- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。该参数中每个Batch的有效token数不超过`query`中的维度S大小且不小于0。支持长度为B的一维tensor。<br>当`layout_query`为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。56+- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。该参数中每个Batch的有效token数不超过`query`中的维度S大小且不小于0。支持长度为B的一维tensor。<br>当`layout_query`为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。
55 57 
56-- **actual\_seq\_lengths\_kv**(`Tensor`):可选参数,表示不同Batch中`key`和`value`的有效token数,数据类型支持`int32`。如果不指定None,表示和key的shape的S长度相同。该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。<br>当`layout_kv`为TND或PA_BSND时,该入参必须传入,`layout_kv`为TND,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。58+- **actual\_seq\_lengths\_kv**(`Tensor`):可选参数,表示不同Batch中`key`和`value`的有效token数,数据类型支持`int32`。如果不指定None,表示和key的shape的S长度相同。该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。<br>当`layout_kv`为TND或PA_BSND时,该入参必须传入,`layout_kv`为TND,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。
57 59 
58-- **sparse\_block\_size**(`int`):可选参数,代表sparse阶段的block大小,在计算importance score时使用,数据类型支持`int64`,支持范围为[1, 16],且为2的幂次方。60+- **sparse\_block\_size**(`int`):可选参数,代表sparse阶段的block大小,在计算importance score时使用,数据类型支持`int64`,支持范围为[1, 16],且为2的幂次方。
59 61 
60-- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,默认值"BSND",支持传入BSND和TND。62+- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,默认值"BSND",支持传入BSND和TND。
61 63 
62-- **layout\_kv**(`str`):可选参数,用于标识输入`key`的数据排布格式,默认值"BSND",支持传入BSND、TND和PA\_BSND,PA\_BSND在使能PageAttention时使用。64+- **layout\_kv**(`str`):可选参数,用于标识输入`key`的数据排布格式,默认值"BSND",支持传入BSND、TND和PA\_BSND,PA\_BSND在使能PageAttention时使用。
63 65 
64-- **sparse\_mode**(`int`):可选参数,表示sparse的模式。数据类型支持`int64`。66+- **sparse\_mode**(`int`):可选参数,表示sparse的模式。数据类型支持`int64`。
65- - sparse\_mode为0时,代表全部计算。67+ - sparse\_mode为0时,代表全部计算。
66- - sparse\_mode为3时,代表rightDownCausal模式的mask,对应以右下顶点往左上为划分线的下三角场景。68+ - sparse\_mode为3时,代表rightDownCausal模式的mask,对应以右下顶点往左上为划分线的下三角场景。
67 69 
68-- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。70+- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。
69 71 
70-- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。72+- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。
71 73 
72-- **attention\_mode**(`int`):可选参数,表示attention的模式。数据类型支持`int64`,仅支持传入2,表示MLA-absorb模式,即QK的D包含rope和nope两部分,且KV是同一份,默认值为0。74+- **attention\_mode**(`int`):可选参数,表示attention的模式。数据类型支持`int64`,仅支持传入2,表示MLA-absorb模式,即QK的D包含rope和nope两部分,且KV是同一份,默认值为0。
73 75 
74-- **quant\_scale\_repo\_mode**(`int`):可选参数,表示量化参数的存放模式。数据类型支持`int64`,仅支持传入1,表示combine模式,即量化参数和数据混合存放,默认值1。76+- **quant\_scale\_repo\_mode**(`int`):可选参数,表示量化参数的存放模式。数据类型支持`int64`,仅支持传入1,表示combine模式,即量化参数和数据混合存放,默认值1。
75 77 
76-- **tile\_size**(`int`):可选参数,表示per_tile时每个参数对应的数据块大小,仅在per_tile时有效。数据类型支持`int64`,仅支持默认值128。78+- **tile\_size**(`int`):可选参数,表示per_tile时每个参数对应的数据块大小,仅在per_tile时有效。数据类型支持`int64`,仅支持默认值128。
77 79 
78-- **rope\_head\_dim**(`int`):可选参数,表示MLA架构下的rope\_head\_dim大小,仅在attention\_mode为2时有效。数据类型支持`int64`,仅支持默认值64。80+- **rope\_head\_dim**(`int`):可选参数,表示MLA架构下的rope\_head\_dim大小,仅在attention\_mode为2时有效。数据类型支持`int64`,仅支持默认值64。
79 81 
80## 返回值说明82## 返回值说明
83+ 
81`Tensor`84`Tensor`
82 85 
83代表公式中的输出Attention。数据格式支持$ND$,数据类型支持`bfloat16``float16`。输出shape与入参`query`的shape保持一致。86代表公式中的输出Attention。数据格式支持$ND$,数据类型支持`bfloat16``float16`。输出shape与入参`query`的shape保持一致。
84 87 
85## 约束说明88## 约束说明
86 89 
87-- 该接口支持推理场景下使用。90+- 该接口支持推理场景下使用。
88-- 该接口支持图模式。91+- 该接口支持图模式。
89-- 参数query中的D值为576,即nope\+rope=512\+64。92+- 参数query中的D值为576,即nope\+rope=512\+64。
90-- 参数key、value中的D值为656,即nope\+rope\*2\+dequant\_scale\*4=512\+64\*2\+4\*4。93+- 参数key、value中的D值为656,即nope\+rope\*2\+dequant\_scale\*4=512\+64\*2\+4\*4。
91-- 支持sparse\_block\_size整除block\_size。94+- 支持sparse\_block\_size整除block\_size。
92-- 非PageAttention场景layout\_query和layout\_kv需要保持一致。95+- 非PageAttention场景layout\_query和layout\_kv需要保持一致。
93 96 
94## 调用示例97## 调用示例
95 98 
96-- 单算子模式调用99+- 单算子模式调用
97 100 
98 ```python101 ```python
99 import torch102 import torch
@@ -153,7 +156,7 @@ torch_npu.npu_kv_quant_sparse_flash_attention(query, key, value, sparse_indices,
153 device='npu:0', dtype=torch.bfloat16)156 device='npu:0', dtype=torch.bfloat16)
154 ```157 ```
155 158 
156-- 图模式调用159+- 图模式调用
157 160 
158 ```python161 ```python
159 import torch162 import torch
@@ -254,4 +257,3 @@ torch_npu.npu_kv_quant_sparse_flash_attention(query, key, value, sparse_indices,
254 [ -256.0000, 256.0000, 596.0000, ..., 92.0000, -736.0000, 0.0000]]]],257 [ -256.0000, 256.0000, 596.0000, ..., 92.0000, -736.0000, 0.0000]]]],
255 device='npu:0', dtype=torch.bfloat16) torch.Size([1, 1, 128, 512])258 device='npu:0', dtype=torch.bfloat16) torch.Size([1, 1, 128, 512])
256 ```259 ```
257- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_kv_rmsnorm_rope_cache.md+54-54
@@ -9,38 +9,38 @@
9 9 
10## 功能说明<a name="zh-cn_topic_0000002236535552_section1023311522369"></a>10## 功能说明<a name="zh-cn_topic_0000002236535552_section1023311522369"></a>
11 11 
12-- API功能:融合了MLA(Multi-head Latent Attention)结构中RMSNorm归一化计算与RoPE(Rotary Position Embedding)位置编码以及更新KVCache的ScatterUpdate操作。12+- API功能:融合了MLA(Multi-head Latent Attention)结构中RMSNorm归一化计算与RoPE(Rotary Position Embedding)位置编码以及更新KVCache的ScatterUpdate操作。
13-- 计算公式:13+- 计算公式:
14- - **输入张量kv拆分**:拆分为两部分,其中B为批次大小,T为序列长度。14+ - **输入张量kv拆分**:拆分为两部分,其中B为批次大小,T为序列长度。
15 15 
16 ![](../../figures/zh-cn_formulaimage_0000002239561238.png)16 ![](../../figures/zh-cn_formulaimage_0000002239561238.png)
17 17 
18- - **RMS归一化**:对rms\_in,应用RMS归一化。18+ - **RMS归一化**:对rms\_in,应用RMS归一化。
19 19 
20 ![](../../figures/zh-cn_formulaimage_0000002239721038.png)20 ![](../../figures/zh-cn_formulaimage_0000002239721038.png)
21 21 
22- - γ∈R^512是可学习的缩放参数。22+ - γ∈R^512是可学习的缩放参数。
23- - Ed\[·\]表示沿最后一个维度(维度d=512)的均值。23+ - Ed\[·\]表示沿最后一个维度(维度d=512)的均值。
24- - ε为小常数(如0.00001),防止除以零。24+ - ε为小常数(如0.00001),防止除以零。
25- - ⊙表示逐元素相乘。25+ - ⊙表示逐元素相乘。
26 26 
27- - **旋转位置编码(RoPE)**27+ - **旋转位置编码(RoPE)**
28- 1. 重塑与转置:将rope\_in重塑并转置以准备旋转。28+ 1. 重塑与转置:将rope\_in重塑并转置以准备旋转。
29 29 
30 ![](../../figures/zh-cn_formulaimage_0000002239561242.png)30 ![](../../figures/zh-cn_formulaimage_0000002239561242.png)
31 31 
32- 2. 旋转操作:应用旋转位置编码。32+ 2. 旋转操作:应用旋转位置编码。
33 33 
34 ![](../../figures/zh-cn_formulaimage_0000002239721042.png)34 ![](../../figures/zh-cn_formulaimage_0000002239721042.png)
35 35 
36- - cos⁡和sin⁡为预计算的旋转角度参数。36+ - cos⁡和sin⁡为预计算的旋转角度参数。
37- - RotateHalf\(k\)将k的后半部分元素移至前半部分并取反,后半部分用前半部分的值。具体来说,对于维度d=64:37+ - RotateHalf\(k\)将k的后半部分元素移至前半部分并取反,后半部分用前半部分的值。具体来说,对于维度d=64:
38 38 
39 ![](../../figures/zh-cn_formulaimage_0000002242091560.png)39 ![](../../figures/zh-cn_formulaimage_0000002242091560.png)
40 40 
41## 函数原型<a name="zh-cn_topic_0000002236535552_section123412524369"></a>41## 函数原型<a name="zh-cn_topic_0000002236535552_section123412524369"></a>
42 42 
43-```43+```python
44torch_npu.npu_kv_rmsnorm_rope_cache(kv, gamma, cos, sin, index, k_cache, ckv_cache, *, k_rope_scale=None, c_kv_scale=None, k_rope_offset=None, c_kv_offset=None, epsilon=1e-5, cache_mode='Norm', is_output_kv=False) -> (Tensor, Tensor, Tensor, Tensor)44torch_npu.npu_kv_rmsnorm_rope_cache(kv, gamma, cos, sin, index, k_cache, ckv_cache, *, k_rope_scale=None, c_kv_scale=None, k_rope_offset=None, c_kv_offset=None, epsilon=1e-5, cache_mode='Norm', is_output_kv=False) -> (Tensor, Tensor, Tensor, Tensor)
45```45```
46 46 
@@ -48,29 +48,30 @@ torch_npu.npu_kv_rmsnorm_rope_cache(kv, gamma, cos, sin, index, k_cache, ckv_cac
48 48 
49> [!NOTE] 49> [!NOTE]
50> Tensor中shape使用的变量说明:50> Tensor中shape使用的变量说明:
51-> - batch\_size:batch的大小。51+>
52-> - seq\_lensequence长度52+> - batch\_sizebatch大小
53-> - hidden\_size表示MLA输入向量长度,取值仅支持57653+> - seq\_lensequence的长度。
54-> - rms\_size:表示RMSNorm分支的向量长度,取值仅支持51254+> - hidden\_size:表示MLA输入的向量长度,取值仅支持576
55-> - rope\_size:表示RoPE分支的向量长度,取值仅支持6455+> - rms\_size:表示RMSNorm分支的向量长度,取值仅支持512
56-> - cache\_lengthNorm模式下有效,表示KVCache最大长度。56+> - rope\_size:表示RoPE分支的向量长度,取值仅支持64
57-> - block\_numPagedAttention模式下有效,表示Block个数57+> - cache\_lengthNorm模式下有效,表示KVCache支持最大长度
58-> - block\_size:PagedAttention模式下有效,表示Block的大小58+> - block\_num:PagedAttention模式下有效,表示Block的个数
59+> - block\_size:PagedAttention模式下有效,表示Block的大小。
59 60 
60-- **kv** (`Tensor`):必选参数,表示输入的特征张量。数据类型支持`bfloat16`、`float16`,数据格式为$BNSD$,要求为4维张量,形状为\[batch\_size, 1, seq\_len, hidden\_size\],其中hidden\_size=rms\_size\(RMS\)+rope\_size\(RoPE\)。61+- **kv** (`Tensor`):必选参数,表示输入的特征张量。数据类型支持`bfloat16`、`float16`,数据格式为$BNSD$,要求为4维张量,形状为\[batch\_size, 1, seq\_len, hidden\_size\],其中hidden\_size=rms\_size\(RMS\)+rope\_size\(RoPE\)。
61-- **gamma** (`Tensor`):必选参数,表示RMS归一化的缩放参数。数据类型支持`bfloat16`、`float16`,数据格式为$ND$,要求为1维张量,形状为\[rms\_size\]。62+- **gamma** (`Tensor`):必选参数,表示RMS归一化的缩放参数。数据类型支持`bfloat16`、`float16`,数据格式为$ND$,要求为1维张量,形状为\[rms\_size\]。
62-- **cos** (`Tensor`):必选参数,表示RoPE旋转位置编码的余弦分量。数据类型支持`bfloat16`、`float16`,数据格式为$ND$,要求为4维张量,形状为\[batch\_size, 1, seq\_len, rope\_size\]。63+- **cos** (`Tensor`):必选参数,表示RoPE旋转位置编码的余弦分量。数据类型支持`bfloat16`、`float16`,数据格式为$ND$,要求为4维张量,形状为\[batch\_size, 1, seq\_len, rope\_size\]。
63-- **sin** (`Tensor`):必选参数,表示RoPE旋转位置编码的正弦分量。数据类型支持`bfloat16`、`float16`,数据格式为$ND$,要求为4维张量,形状为\[batch\_size, 1, seq\_len, rope\_size\]。64+- **sin** (`Tensor`):必选参数,表示RoPE旋转位置编码的正弦分量。数据类型支持`bfloat16`、`float16`,数据格式为$ND$,要求为4维张量,形状为\[batch\_size, 1, seq\_len, rope\_size\]。
64-- **index** (`Tensor`):必选参数,表示缓存索引张量,用于定位`k_cache`和`ckv_cache`的写入位置。数据类型支持`int64`,数据格式为$ND$。shape取决于`cache_mode`。65+- **index** (`Tensor`):必选参数,表示缓存索引张量,用于定位`k_cache`和`ckv_cache`的写入位置。数据类型支持`int64`,数据格式为$ND$。shape取决于`cache_mode`。
65-- **k\_cache** (`Tensor`):必选参数,用于存储量化/非量化的键向量。数据类型支持`bfloat16`、`float16`、`int8`,数据格式为$ND$。shape取决于`cache_mode`。66+- **k\_cache** (`Tensor`):必选参数,用于存储量化/非量化的键向量。数据类型支持`bfloat16`、`float16`、`int8`,数据格式为$ND$。shape取决于`cache_mode`。
66-- **ckv\_cache** (`Tensor`):必选参数,用于存储量化/非量化的压缩后的kv向量。数据类型支持`bfloat16`、`float16`、`int8`,数据格式为$ND$。shape取决于`cache_mode`。67+- **ckv\_cache** (`Tensor`):必选参数,用于存储量化/非量化的压缩后的kv向量。数据类型支持`bfloat16`、`float16`、`int8`,数据格式为$ND$。shape取决于`cache_mode`。
67- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。68- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
68-- **k\_rope\_scale** (`Tensor`):可选参数,默认值None,表示k旋转位置编码的量化缩放因子。数据类型支持`float32`,数据格式为$ND$,要求为1维张量,形状为\[rope\_size\]。量化模式下必填。69+- **k\_rope\_scale** (`Tensor`):可选参数,默认值None,表示k旋转位置编码的量化缩放因子。数据类型支持`float32`,数据格式为$ND$,要求为1维张量,形状为\[rope\_size\]。量化模式下必填。
69-- **c\_kv\_scale** (`Tensor`):可选参数,默认值None,表示压缩后kv的量化缩放因子。数据类型支持`float32`,数据格式为$ND$,要求为1维张量,形状为\[rms\_size\]。量化模式下必填。70+- **c\_kv\_scale** (`Tensor`):可选参数,默认值None,表示压缩后kv的量化缩放因子。数据类型支持`float32`,数据格式为$ND$,要求为1维张量,形状为\[rms\_size\]。量化模式下必填。
70-- **k\_rope\_offset** (`Tensor`):可选参数,默认值None,表示k旋转位置编码量化偏移量。数据类型支持`float32`,数据格式为$ND$,要求为1维张量,形状为\[rope\_size\]。量化模式下必填。71+- **k\_rope\_offset** (`Tensor`):可选参数,默认值None,表示k旋转位置编码量化偏移量。数据类型支持`float32`,数据格式为$ND$,要求为1维张量,形状为\[rope\_size\]。量化模式下必填。
71-- **c\_kv\_offset** (`Tensor`):可选参数,默认值None,表示压缩后kv的量化偏移量。数据类型支持`float32`,数据格式为$ND$,要求为1维张量,形状为\[rms\_size\]。量化模式下必填。72+- **c\_kv\_offset** (`Tensor`):可选参数,默认值None,表示压缩后kv的量化偏移量。数据类型支持`float32`,数据格式为$ND$,要求为1维张量,形状为\[rms\_size\]。量化模式下必填。
72-- **epsilon** (`float`):可选参数,默认值1e-5,表示RMS归一化中的极小值,防止除以零。73+- **epsilon** (`float`):可选参数,默认值1e-5,表示RMS归一化中的极小值,防止除以零。
73-- **cache\_mode** (`str`):可选参数,默认值'Norm',表示缓存模式,支持的模式如下:74+- **cache\_mode** (`str`):可选参数,默认值'Norm',表示缓存模式,支持的模式如下:
74 75 
75 <a name="zh-cn_topic_0000002236535552_table16997195773911"></a>76 <a name="zh-cn_topic_0000002236535552_table16997195773911"></a>
76 <table><thead align="left"><tr id="zh-cn_topic_0000002236535552_row12998195743918"><th class="cellrowborder" valign="top" width="10.34%" id="mcps1.1.4.1.1"><p id="zh-cn_topic_0000002236535552_p1299819576394"><a name="zh-cn_topic_0000002236535552_p1299819576394"></a><a name="zh-cn_topic_0000002236535552_p1299819576394"></a>枚举值</p>77 <table><thead align="left"><tr id="zh-cn_topic_0000002236535552_row12998195743918"><th class="cellrowborder" valign="top" width="10.34%" id="mcps1.1.4.1.1"><p id="zh-cn_topic_0000002236535552_p1299819576394"><a name="zh-cn_topic_0000002236535552_p1299819576394"></a><a name="zh-cn_topic_0000002236535552_p1299819576394"></a>枚举值</p>
@@ -123,33 +124,33 @@ torch_npu.npu_kv_rmsnorm_rope_cache(kv, gamma, cos, sin, index, k_cache, ckv_cac
123 </tbody>124 </tbody>
124 </table>125 </table>
125 126 
126-- **is\_output\_kv** (`bool`):可选参数,表示是否输出处理后的`k_embed_out`和`y_out`(未量化的原始值),默认值False不输出,仅`cache_mode`在\(PA/PA\_BNSD/PA\_NZ/PA\_BLK\_BNSD/PA\_BLK\_NZ\)模式下有效。127+- **is\_output\_kv** (`bool`):可选参数,表示是否输出处理后的`k_embed_out`和`y_out`(未量化的原始值),默认值False不输出,仅`cache_mode`在\(PA/PA\_BNSD/PA\_NZ/PA\_BLK\_BNSD/PA\_BLK\_NZ\)模式下有效。
127 128 
128## 返回值说明<a name="zh-cn_topic_0000002236535552_section3234185215368"></a>129## 返回值说明<a name="zh-cn_topic_0000002236535552_section3234185215368"></a>
129 130 
130-- **k\_cache** (`Tensor`):和输入`k_cache`的数据类型、维度、数据格式完全一致(本质in-place更新)。131+- **k\_cache** (`Tensor`):和输入`k_cache`的数据类型、维度、数据格式完全一致(本质in-place更新)。
131-- **ckv\_cache** (`Tensor`):和输入`ckv_cache`的数据类型、维度、数据格式完全一致(本质in-place更新)。132+- **ckv\_cache** (`Tensor`):和输入`ckv_cache`的数据类型、维度、数据格式完全一致(本质in-place更新)。
132-- **k\_embed\_out** (`Tensor`):仅当`is_output_kv`为True时,表示RoPE处理后的值。要求为4维张量,形状为\[batch\_size, 1, seq\_len, 64\],数据类型和格式同输入`kv`一致。133+- **k\_embed\_out** (`Tensor`):仅当`is_output_kv`为True时,表示RoPE处理后的值。要求为4维张量,形状为\[batch\_size, 1, seq\_len, 64\],数据类型和格式同输入`kv`一致。
133-- **y\_out** (`Tensor`):仅当`is_output_kv`为True时,表示RMSNorm处理后的值。要求为4维张量,形状为\[batch\_size, 1, seq\_len, 512\],数据类型和格式同输入`kv`一致。134+- **y\_out** (`Tensor`):仅当`is_output_kv`为True时,表示RMSNorm处理后的值。要求为4维张量,形状为\[batch\_size, 1, seq\_len, 512\],数据类型和格式同输入`kv`一致。
134 135 
135## 约束说明<a name="zh-cn_topic_0000002236535552_section1523425283618"></a>136## 约束说明<a name="zh-cn_topic_0000002236535552_section1523425283618"></a>
136 137 
137-- 该接口支持推理场景下使用。138+- 该接口支持推理场景下使用。
138-- 该接口支持图模式。139+- 该接口支持图模式。
139-- 量化模式:当`k_rope_scale`和`c_kv_scale`非空时,`k_cache`和`ckv_cache`的dtype为`int8`,缓存形状的最后一个维度需要为32(Cache数据格式为FRACTAL\_NZ模式),`k_rope_scale`和`c_kv_scale`必须同时非空,`k_rope_offset`和`c_kv_offset`必须同时为None或非空。140+- 量化模式:当`k_rope_scale`和`c_kv_scale`非空时,`k_cache`和`ckv_cache`的dtype为`int8`,缓存形状的最后一个维度需要为32(Cache数据格式为FRACTAL\_NZ模式),`k_rope_scale`和`c_kv_scale`必须同时非空,`k_rope_offset`和`c_kv_offset`必须同时为None或非空。
140-- 非量化模式:当`k_rope_scale`和`c_kv_scale`为空时,`k_cache`和`ckv_cache`的dtype为`bfloat16`或`float16`。141+- 非量化模式:当`k_rope_scale`和`c_kv_scale`为空时,`k_cache`和`ckv_cache`的dtype为`bfloat16`或`float16`。
141-- 索引映射:所有`cache_mode`缓存模式下,index的值不可以重复,如果传入的index值存在重复,算子的行为是未定义的且不可预知的。142+- 索引映射:所有`cache_mode`缓存模式下,index的值不可以重复,如果传入的index值存在重复,算子的行为是未定义的且不可预知的。
142- - Norm:index的值表示每个Batch下的偏移。143+ - Norm:index的值表示每个Batch下的偏移。
143- - PA/PA\_BNSD/PA\_NZ:index的值表示全局的偏移。144+ - PA/PA\_BNSD/PA\_NZ:index的值表示全局的偏移。
144- - PA\_BLK\_BNSD/PA\_BLK\_NZ:index的值表示每个页的全局偏移;这个场景假设cache更新是连续的,不支持非连续更新的cache。145+ - PA\_BLK\_BNSD/PA\_BLK\_NZ:index的值表示每个页的全局偏移;这个场景假设cache更新是连续的,不支持非连续更新的cache。
145 146 
146-- Shape关联规则:不同的`cache_mode`缓存模式有不同的Shape规则。147+- Shape关联规则:不同的`cache_mode`缓存模式有不同的Shape规则。
147- - Norm:k\_cache形状为\[batch\_size, 1, cache\_length, rope\_size\],ckv\_cache形状为\[batch\_size, 1, cache\_length, rms\_size\],index形状为\[batch\_size, seq\_len\], cache\_length\>=seq\_len。148+ - Norm:k\_cache形状为\[batch\_size, 1, cache\_length, rope\_size\],ckv\_cache形状为\[batch\_size, 1, cache\_length, rms\_size\],index形状为\[batch\_size, seq\_len\], cache\_length\>=seq\_len。
148- - 非Norm模式\(PagedAttention相关模式\):要求block\_num\>=Ceil\(seq\_len/block\_size\)\*batch\_size。149+ - 非Norm模式\(PagedAttention相关模式\):要求block\_num\>=Ceil\(seq\_len/block\_size\)\*batch\_size。
149 150 
150## 调用示例<a name="zh-cn_topic_0000002236535552_section3235105212365"></a>151## 调用示例<a name="zh-cn_topic_0000002236535552_section3235105212365"></a>
151 152 
152-- 单算子模式调用153+- 单算子模式调用
153 154 
154 ```python155 ```python
155 import torch156 import torch
@@ -199,7 +200,7 @@ torch_npu.npu_kv_rmsnorm_rope_cache(kv, gamma, cos, sin, index, k_cache, ckv_cac
199 200
200 ```201 ```
201 202 
202-- 图模式调用203+- 图模式调用
203 204 
204 ```python205 ```python
205 import torch206 import torch
@@ -253,4 +254,3 @@ torch_npu.npu_kv_rmsnorm_rope_cache(kv, gamma, cos, sin, index, k_cache, ckv_cac
253 model = torch.compile(model, backend=npu_backend, dynamic=False)254 model = torch.compile(model, backend=npu_backend, dynamic=False)
254 _, _, k_rope, c_kv = model(kv, gamma, cos, sin, index, k_cache, ckv_cache, k_rope_scale, c_kv_scale, None, None, 1e-5, cache_mode, is_output_kv)255 _, _, k_rope, c_kv = model(kv, gamma, cos, sin, index, k_cache, ckv_cache, k_rope_scale, c_kv_scale, None, None, 1e-5, cache_mode, is_output_kv)
255 ```256 ```
256- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_lightning_indexer.md+31-31
@@ -1,6 +1,7 @@
1# torch_npu-npu_lightning_indexer1# torch_npu-npu_lightning_indexer
2 2 
3## 产品支持情况3## 产品支持情况
4+ 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
5| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
6|<term>Atlas A2 推理系列产品</term> | √ |7|<term>Atlas A2 推理系列产品</term> | √ |
@@ -8,9 +9,9 @@
8 9 
9## 功能说明10## 功能说明
10 11 
11-- API功能:`lightning_indexer`基于一系列操作得到每一个token对应的Top-$k$个位置。12+- API功能:`lightning_indexer`基于一系列操作得到每一个token对应的Top-$k$个位置。
12 13 
13-- 计算公式:14+- 计算公式:
14 $$15 $$
15 Indices=\text{Top-}k\left\{[1]_{1\times g}@\left[(W@[1]_{1\times S_{k}})\odot\text{ReLU}\left(Q_{index}@K_{index}^T\right)\right]\right\}16 Indices=\text{Top-}k\left\{[1]_{1\times g}@\left[(W@[1]_{1\times S_{k}})\odot\text{ReLU}\left(Q_{index}@K_{index}^T\right)\right]\right\}
16 $$17 $$
@@ -18,7 +19,7 @@
18 19 
19## 函数原型20## 函数原型
20 21 
21-```22+```python
22torch_npu.npu_lightning_indexer(query, key, weights, *, actual_seq_lengths_query=None, actual_seq_lengths_key=None, block_table=None, layout_query="BSND", layout_key="BSND", sparse_count=2048, sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1, return_value=False) -> (Tensor, Tensor)23torch_npu.npu_lightning_indexer(query, key, weights, *, actual_seq_lengths_query=None, actual_seq_lengths_key=None, block_table=None, layout_query="BSND", layout_key="BSND", sparse_count=2048, sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1, return_value=False) -> (Tensor, Tensor)
23```24```
24 25 
@@ -29,57 +30,57 @@ torch_npu.npu_lightning_indexer(query, key, weights, *, actual_seq_lengths_query
29> - query、key、weights参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。30> - query、key、weights参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。
30> - S1表示query shape中的S,S2表示key shape中的S,T1表示query shape中的T,T2表示key shape中的T,N1表示query shape中的N,N2表示key shape中的N。31> - S1表示query shape中的S,S2表示key shape中的S,T1表示query shape中的T,T2表示key shape中的T,N1表示query shape中的N,N2表示key shape中的N。
31 32 
32-- **query**(`Tensor`):必选参数,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`和`float16`。`layout_query`为BSND时shape为[B,S1,N1,D],当`layout_query`为TND时shape为[T1,N1,D],N1仅支持小于等于64。33+- **query**(`Tensor`):必选参数,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`和`float16`。`layout_query`为BSND时shape为[B,S1,N1,D],当`layout_query`为TND时shape为[T1,N1,D],N1仅支持小于等于64。
33 34
34-- **key**(`Tensor`):必选参数,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`和`float16`,layout\_key为PA_BSND时shape为[block\_count, block\_size, N2, D],其中block\_count为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的整数倍,最大支持到1024。`layout_key`为BSND时shape为[B, S2, N2, D],`layout_key`为TND时shape为[T2, N2, D],N2仅支持1。35+- **key**(`Tensor`):必选参数,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`和`float16`,layout\_key为PA_BSND时shape为[block\_count, block\_size, N2, D],其中block\_count为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的整数倍,最大支持到1024。`layout_key`为BSND时shape为[B, S2, N2, D],`layout_key`为TND时shape为[T2, N2, D],N2仅支持1。
35 36
36-- **weights**(`Tensor`):必选参数,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`、`float16`和`float32`,支持输入shape[B,S1,N1]、[T,N1]。37+- **weights**(`Tensor`):必选参数,不支持非连续,数据格式支持$ND$,数据类型支持`bfloat16`、`float16`和`float32`,支持输入shape[B,S1,N1]、[T,N1]。
37 38
38- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。39- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
39 40 
40-- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。41+- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。
41- - 该入参中每个Batch的有效token数不超过`query`中的维度S大小且不小于0。支持长度为B的一维tensor。当`layout_query`为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。42+ - 该入参中每个Batch的有效token数不超过`query`中的维度S大小且不小于0。支持长度为B的一维tensor。当`layout_query`为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。
42 43 
43-- **actual\_seq\_lengths\_key**(`Tensor`):可选参数,表示不同Batch中`key`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和key的shape的S长度相同。44+- **actual\_seq\_lengths\_key**(`Tensor`):可选参数,表示不同Batch中`key`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和key的shape的S长度相同。
44- - 该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。当`layout_key`为TND或PA_BSND时,该入参必须传入,`layout_key`为TND,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。45+ - 该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。当`layout_key`为TND或PA_BSND时,该入参必须传入,`layout_key`为TND,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。
45 46 
46-- **block\_table**(`Tensor`):可选参数,表示PageAttention中KV存储使用的block映射表,数据格式支持$ND$,数据类型支持`int32`。47+- **block\_table**(`Tensor`):可选参数,表示PageAttention中KV存储使用的block映射表,数据格式支持$ND$,数据类型支持`int32`。
47- - PageAttention场景下,block\_table必须为二维,第一维长度需要等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为每个batch中最大actual\_seq\_lengths\_key对应的block数量)48+ - PageAttention场景下,block\_table必须为二维,第一维长度需要等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为每个batch中最大actual\_seq\_lengths\_key对应的block数量)
48 49 
49-- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,当前支持BSND、TND,默认值"BSND"。50+- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,当前支持BSND、TND,默认值"BSND"。
50 51 
51-- **layout\_key**(`str`):可选参数,用于标识输入`key`的数据排布格式,当前支持PA_BSND、BSND、TND,默认值"BSND",在非PageAttention场景下,该参数值应与**layout\_query**值保持一致。52+- **layout\_key**(`str`):可选参数,用于标识输入`key`的数据排布格式,当前支持PA_BSND、BSND、TND,默认值"BSND",在非PageAttention场景下,该参数值应与**layout\_query**值保持一致。
52 53 
53-- **sparse\_count**(`int`):可选参数,代表topK阶段需要保留的block数量,支持[1, 2048]以及3072、4096、5120、6144、7168、8192,数据类型支持`int32`。54+- **sparse\_count**(`int`):可选参数,代表topK阶段需要保留的block数量,支持[1, 2048]以及3072、4096、5120、6144、7168、8192,数据类型支持`int32`。
54 55 
55-- **sparse\_mode**(`int`):可选参数,表示sparse的模式,支持0/3,数据类型支持`int32`。56+- **sparse\_mode**(`int`):可选参数,表示sparse的模式,支持0/3,数据类型支持`int32`。
56 57
57- - sparse\_mode为0时,代表defaultMask模式。58+ - sparse\_mode为0时,代表defaultMask模式。
58- - sparse\_mode为3时,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景。59+ - sparse\_mode为3时,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景。
59 60 
60-- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`。仅支持默认值2^63-1。61+- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`。仅支持默认值2^63-1。
61 62 
62-- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`。仅支持默认值2^63-1。63+- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`。仅支持默认值2^63-1。
63 64 
64-- **return\_value**(`bool`):可选参数,表示是否输出`sparse_values`。True表示输出,但图模式下不支持,False表示不输出;默认值为False。该参数仅在训练且`layout_key`不为PA_BSND场景支持。65+- **return\_value**(`bool`):可选参数,表示是否输出`sparse_values`。True表示输出,但图模式下不支持,False表示不输出;默认值为False。该参数仅在训练且`layout_key`不为PA_BSND场景支持。
65 66 
66## 返回值说明67## 返回值说明
67 68 
68-- **sparse\_indices**(`Tensor`):公式中的Indices输出,数据类型支持`int32`,数据格式支持$ND$,当`layout_query`为"BSND"时输出shape为[B, S1, N2, sparse\_count],当layout\_query为"TND"时输出shape为[T1, N2, sparse\_count]。69+- **sparse\_indices**(`Tensor`):公式中的Indices输出,数据类型支持`int32`,数据格式支持$ND$,当`layout_query`为"BSND"时输出shape为[B, S1, N2, sparse\_count],当layout\_query为"TND"时输出shape为[T1, N2, sparse\_count]。
69 70 
70-- **sparse\_values**(`Tensor`):公式中的Indices输出对应的value值,数据类型支持`bfloat16`、`float16`,数据格式支持$ND$,输出shape与`sparse_indices`保持一致。71+- **sparse\_values**(`Tensor`):公式中的Indices输出对应的value值,数据类型支持`bfloat16`、`float16`,数据格式支持$ND$,输出shape与`sparse_indices`保持一致。
71 72 
72## 约束说明73## 约束说明
73 74 
74-- 该接口支持图模式。75+- 该接口支持图模式。
75-- 参数query中的N支持小于等于64,key中的N支持1。76+- 参数query中的N支持小于等于64,key中的N支持1。
76-- 参数query中的D和参数key中的D值相等为128。77+- 参数query中的D和参数key中的D值相等为128。
77-- 参数query、key的数据类型应保持一致。78+- 参数query、key的数据类型应保持一致。
78-- 参数weights不为`float32`时,参数query、key、weights的数据类型应保持一致。79+- 参数weights不为`float32`时,参数query、key、weights的数据类型应保持一致。
79 80 
80## 调用示例81## 调用示例
81 82 
82-- 单算子模式调用83+- 单算子模式调用
83 84 
84 ```python85 ```python
85 import torch86 import torch
@@ -117,7 +118,7 @@ torch_npu.npu_lightning_indexer(query, key, weights, *, actual_seq_lengths_query
117 device='npu:0', dtype=torch.int32)118 device='npu:0', dtype=torch.int32)
118 ```119 ```
119 120 
120-- 图模式调用121+- 图模式调用
121 122 
122 ```python123 ```python
123 import torch124 import torch
@@ -187,4 +188,3 @@ torch_npu.npu_lightning_indexer(query, key, weights, *, actual_seq_lengths_query
187 graph output: tensor([[[[4488, 3926, 1154, ..., 3535, 8031, 8180]]]],188 graph output: tensor([[[[4488, 3926, 1154, ..., 3535, 8031, 8180]]]],
188 device='npu:0', dtype=torch.int32) torch.Size([1, 1, 1, 2048])189 device='npu:0', dtype=torch.int32) torch.Size([1, 1, 1, 2048])
189 ```190 ```
190- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_matmul_all_to_all.md+18-18
@@ -20,37 +20,37 @@
20 20 
21## 函数原型21## 函数原型
22 22 
23-```23+```python
24torch_npu.npu_matmul_all_to_all(x1, x2, hcom, world_size, bias=None, all2all_axes=None) -> Tensor24torch_npu.npu_matmul_all_to_all(x1, x2, hcom, world_size, bias=None, all2all_axes=None) -> Tensor
25```25```
26 26 
27## 参数说明27## 参数说明
28 28 
29-- **x1**(`Tensor`):必选输入,表示融合算子的左矩阵输入,也是Matmul计算的左矩阵,对应公式中的x1。数据类型支持bfloat16、float16,维度只能为2D,shape为(BS, H1),数据格式支持ND,不支持非连续Tensor,支持第一维度为0的空Tensor。29+- **x1**(`Tensor`):必选输入,表示融合算子的左矩阵输入,也是Matmul计算的左矩阵,对应公式中的x1。数据类型支持bfloat16、float16,维度只能为2D,shape为(BS, H1),数据格式支持ND,不支持非连续Tensor,支持第一维度为0的空Tensor。
30-- **x2**(`Tensor`):必选输入,表示融合算子的右矩阵输入,也是Matmul计算的右矩阵,对应公式中的x2。数据类型与x1一致,维度只能为2D,shape为(H1, H2),数据格式支持ND,支持转置非连续Tensor。30+- **x2**(`Tensor`):必选输入,表示融合算子的右矩阵输入,也是Matmul计算的右矩阵,对应公式中的x2。数据类型与x1一致,维度只能为2D,shape为(H1, H2),数据格式支持ND,支持转置非连续Tensor。
31-- **hcom**(`str`):必选输入,Host侧标识列组的字符串,即通信域名称,通过get_hccl_comm_name接口获取。31+- **hcom**(`str`):必选输入,Host侧标识列组的字符串,即通信域名称,通过get_hccl_comm_name接口获取。
32-- **world_size**(`int`):必选输入,通信域内的rank总数,对应公式中的rankSize,支持范围[2, 4, 8, 16]。32+- **world_size**(`int`):必选输入,通信域内的rank总数,对应公式中的rankSize,支持范围[2, 4, 8, 16]。
33-- **bias**(`Tensor`):可选输入,矩阵乘运算后累加的偏置,对应公式中的bias。数据类型由输入x1和x2决定,当x1和x2为float16时,bias的数据类型为float16;当x1和x2为bfloat16时,bias的数据类型为float32。维度只能为1D,shape为(H2),数据类型支持ND。33+- **bias**(`Tensor`):可选输入,矩阵乘运算后累加的偏置,对应公式中的bias。数据类型由输入x1和x2决定,当x1和x2为float16时,bias的数据类型为float16;当x1和x2为bfloat16时,bias的数据类型为float32。维度只能为1D,shape为(H2),数据类型支持ND。
34-- **all2all_axes**(`List[int]`):可选输入,AlltoAll和Permute数据交换的方向,支持为空或者[-1, -2],表示将Matmul结果由(BS, H2)转为(BS*rankSize, H2/rankSize)。34+- **all2all_axes**(`List[int]`):可选输入,AlltoAll和Permute数据交换的方向,支持为空或者[-1, -2],表示将Matmul结果由(BS, H2)转为(BS*rankSize, H2/rankSize)。
35 35 
36## 返回值说明36## 返回值说明
37 37 
38-- **y**(`Tensor`):计算输出,表示最终的计算结果output,数据类型与输入x1或者x2保持一致,支持2维,shape为(BS*rankSize, H2/rankSize),数据格式支持ND,不支持非连续的Tensor。38+- **y**(`Tensor`):计算输出,表示最终的计算结果output,数据类型与输入x1或者x2保持一致,支持2维,shape为(BS*rankSize, H2/rankSize),数据格式支持ND,不支持非连续的Tensor。
39 39 
40## 约束说明40## 约束说明
41 41 
42-- 该接口支持训练、推理场景下使用。42+- 该接口支持训练、推理场景下使用。
43-- A3场景下,该接口支持单算子模式,不支持图模式。43+- A3场景下,该接口支持单算子模式,不支持图模式。
44-- 除x1以外的输入参数均不支持空Tensor。44+- 除x1以外的输入参数均不支持空Tensor。
45-- 通信域名称hcom不支持传入空字符串,长度取值范围为[1, 127]。45+- 通信域名称hcom不支持传入空字符串,长度取值范围为[1, 127]。
46-- 输入参数Tensor中shape使用的变量说明:46+- 输入参数Tensor中shape使用的变量说明:
47- - BS:输入左矩阵的第一维度大小,表示输入序列sequence的条数,BS*rankSize取值范围为[0, 2147483647]。47+ - BS:输入左矩阵的第一维度大小,表示输入序列sequence的条数,BS*rankSize取值范围为[0, 2147483647]。
48- - H1:输入左矩阵的第二维度大小和输入右矩阵的第一维度大小,表示隐藏层维度,取值范围为[1, 65535]。48+ - H1:输入左矩阵的第二维度大小和输入右矩阵的第一维度大小,表示隐藏层维度,取值范围为[1, 65535]。
49- - H2:输入右矩阵的第二维度大小,表示输出序列sequence的长度,取值范围为[2, 2147483647],必须整除rankSize。49+ - H2:输入右矩阵的第二维度大小,表示输出序列sequence的长度,取值范围为[2, 2147483647],必须整除rankSize。
50 50 
51## 调用示例51## 调用示例
52 52 
53-- 单算子模式调用53+- 单算子模式调用
54 54 
55 ```python55 ```python
56 import torch56 import torch
@@ -89,4 +89,4 @@ torch_npu.npu_matmul_all_to_all(x1, x2, hcom, world_size, bias=None, all2all_axe
89 args=(worksize, master_ip, master_port, x1_shape, x2_shape),89 args=(worksize, master_ip, master_port, x1_shape, x2_shape),
90 nprocs=worksize,90 nprocs=worksize,
91 )91 )
92- ```92+ ```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_mla_prolog.md+34-36
@@ -7,7 +7,6 @@
7|<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |7|<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |
8|<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |8|<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |
9 9 
10- 
11## 功能说明10## 功能说明
12 11 
13- API功能:推理场景下,Multi-Head Latent Attention(MLA)前处理的计算。12- API功能:推理场景下,Multi-Head Latent Attention(MLA)前处理的计算。
@@ -18,27 +17,27 @@
18 - 第三路是输入x乘以W<sup>DKV</sup>进行下采样和RmsNorm后传入Cache中得到k<sup>C</sup>17 - 第三路是输入x乘以W<sup>DKV</sup>进行下采样和RmsNorm后传入Cache中得到k<sup>C</sup>
19 - 第四路是输入x乘以W<sup>KR</sup>后经过旋转位置编码后传入另一个Cache中得到k<sup>R</sup>18 - 第四路是输入x乘以W<sup>KR</sup>后经过旋转位置编码后传入另一个Cache中得到k<sup>R</sup>
20 19
21-- 计算公式:20+- 计算公式:
22- - RmsNorm公式:21+ - RmsNorm公式:
23 $$RmsNorm(x) = \gamma \cdot \frac{x} {\sqrt{\frac{1}{N} \sum_{i=1}^{N} x_i^2 + \epsilon}}$$22 $$RmsNorm(x) = \gamma \cdot \frac{x} {\sqrt{\frac{1}{N} \sum_{i=1}^{N} x_i^2 + \epsilon}}$$
24 23 
25- - Query计算公式:24+ - Query计算公式:
26 $$c^Q = RmsNorm(x \cdot W^{DQ})$$25 $$c^Q = RmsNorm(x \cdot W^{DQ})$$
27 $$q^C = c^Q \cdot W^{UQ}$$26 $$q^C = c^Q \cdot W^{UQ}$$
28 $$q^N = q^C \cdot W^{UK}$$27 $$q^N = q^C \cdot W^{UK}$$
29 28 
30- - Query ROPE旋转位置编码:29+ - Query ROPE旋转位置编码:
31 $$q^R = ROPE(c^Q \cdot W^{QR})$$30 $$q^R = ROPE(c^Q \cdot W^{QR})$$
32 31 
33- - Key计算公式:32+ - Key计算公式:
34 $$k^C = Cache(RmsNorm(x \cdot W^{DKV}))$$33 $$k^C = Cache(RmsNorm(x \cdot W^{DKV}))$$
35 34 
36- - Key ROPE旋转位置编码:35+ - Key ROPE旋转位置编码:
37 $$k^R = Cache(ROPE(x \cdot W^{KR}))$$36 $$k^R = Cache(ROPE(x \cdot W^{KR}))$$
38 37 
39## 函数原型38## 函数原型
40 39 
41-```40+```python
42torch_npu.npu_mla_prolog(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv_kr, rmsnorm_gamma_cq, rmsnorm_gamma_ckv, rope_sin, rope_cos, cache_index, kv_cache, kr_cache, *, dequant_scale_x=None, dequant_scale_w_dq=None, dequant_scale_w_uq_qr=None, dequant_scale_w_dkv_kr=None, quant_scale_ckv=None, quant_scale_ckr=None, smooth_scales_cq=None, rmsnorm_epsilon_cq=1e-05, rmsnorm_epsilon_ckv=1e-05, cache_mode="PA_BSND") -> (Tensor, Tensor, Tensor, Tensor)41torch_npu.npu_mla_prolog(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv_kr, rmsnorm_gamma_cq, rmsnorm_gamma_ckv, rope_sin, rope_cos, cache_index, kv_cache, kr_cache, *, dequant_scale_x=None, dequant_scale_w_dq=None, dequant_scale_w_uq_qr=None, dequant_scale_w_dkv_kr=None, quant_scale_ckv=None, quant_scale_ckr=None, smooth_scales_cq=None, rmsnorm_epsilon_cq=1e-05, rmsnorm_epsilon_ckv=1e-05, cache_mode="PA_BSND") -> (Tensor, Tensor, Tensor, Tensor)
43```42```
44 43 
@@ -47,10 +46,10 @@ torch_npu.npu_mla_prolog(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv
47- **token\_x**(`Tensor`):必选参数,对应公式中x。shape支持2维和3维,格式为\(T, He\)和\(B, S, He\),dtype支持`bfloat16`,数据格式支持ND。46- **token\_x**(`Tensor`):必选参数,对应公式中x。shape支持2维和3维,格式为\(T, He\)和\(B, S, He\),dtype支持`bfloat16`,数据格式支持ND。
48- **weight\_dq**(`Tensor`):必选参数,表示计算Query的下采样权重矩阵,即公式中W<sup>DQ</sup>。shape支持2维,格式为\(He, Hcq\),dtype支持`bfloat16`,数据格式支持FRACTAL\_NZ(可通过`torch_npu.npu_format_cast`将ND格式转为FRACTAL\_NZ格式)。47- **weight\_dq**(`Tensor`):必选参数,表示计算Query的下采样权重矩阵,即公式中W<sup>DQ</sup>。shape支持2维,格式为\(He, Hcq\),dtype支持`bfloat16`,数据格式支持FRACTAL\_NZ(可通过`torch_npu.npu_format_cast`将ND格式转为FRACTAL\_NZ格式)。
49- **weight\_uq\_qr**`Tensor`):必选参数,表示计算Query的上采样权重矩阵和Query的位置编码权重矩阵,即公式中W<sup>UQ</sup>和W<sup>QR</sup>。shape支持2维,格式为\(Hcq, N\*\(D+Dr\)\),dtype支持`bfloat16`和`int8`,数据格式支持FRACTAL\_NZ。 48- **weight\_uq\_qr**`Tensor`):必选参数,表示计算Query的上采样权重矩阵和Query的位置编码权重矩阵,即公式中W<sup>UQ</sup>和W<sup>QR</sup>。shape支持2维,格式为\(Hcq, N\*\(D+Dr\)\),dtype支持`bfloat16`和`int8`,数据格式支持FRACTAL\_NZ。
50- - 当`weight_uq_qr`为`int8`类型时,`weight_uq_qr`是一个pertensor的量化后的输入,表示当前为部分量化场景。 49+ - 当`weight_uq_qr`为`int8`类型时,`weight_uq_qr`是一个pertensor的量化后的输入,表示当前为部分量化场景。
51 -`kv_cache``kr_cache``bfloat16`类型,对应`kv_cache_out``kr_cache_out`为非量化输出,此时`dequant_scale_w_uq_qr`字段必须传入,`smooth_scales_cq`字段可选传入。 50 -`kv_cache``kr_cache``bfloat16`类型,对应`kv_cache_out``kr_cache_out`为非量化输出,此时`dequant_scale_w_uq_qr`字段必须传入,`smooth_scales_cq`字段可选传入。
52 -`kv_cache``kr_cache``int8`类型,对应`kv_cache_out``kr_cache_out`为量化输出,此时`dequant_scale_w_uq_qr``quant_scale_ckv``quant_scale_ckr`字段必须传入,`smooth_scales_cq`字段可选传入。 51 -`kv_cache``kr_cache``int8`类型,对应`kv_cache_out``kr_cache_out`为量化输出,此时`dequant_scale_w_uq_qr``quant_scale_ckv``quant_scale_ckr`字段必须传入,`smooth_scales_cq`字段可选传入。
53- - 当`weight_uq_qr`为`bfloat16`类型时,表示当前为非量化场景。 52+ - 当`weight_uq_qr`为`bfloat16`类型时,表示当前为非量化场景。
54 此时`dequant_scale_w_uq_qr`、`quant_scale_ckv`、`quant_scale_ckr`、`smooth_scales_cq`字段不能传入(即为none)。53 此时`dequant_scale_w_uq_qr`、`quant_scale_ckv`、`quant_scale_ckr`、`smooth_scales_cq`字段不能传入(即为none)。
55 54
56- **weight\_uk**(`Tensor`):必选参数,表示计算Key的上采样权重,即公式中W<sup>UK</sup>。shape支持3维,格式为\(N, D, Hckv\),dtype支持`bfloat16`,数据格式支持ND。55- **weight\_uk**(`Tensor`):必选参数,表示计算Key的上采样权重,即公式中W<sup>UK</sup>。shape支持3维,格式为\(N, D, Hckv\),dtype支持`bfloat16`,数据格式支持ND。
@@ -76,38 +75,38 @@ torch_npu.npu_mla_prolog(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv
76 75 
77## 返回值说明76## 返回值说明
78 77 
79-- **query**(`Tensor`):表示Query的输出Tensor,即公式中q<sup>N</sup>。shape支持3维和4维,格式为\(T, N, Hckv\)和\(B, S, N, Hckv\),dtype支持`bfloat16`,数据格式支持ND。78+- **query**(`Tensor`):表示Query的输出Tensor,即公式中q<sup>N</sup>。shape支持3维和4维,格式为\(T, N, Hckv\)和\(B, S, N, Hckv\),dtype支持`bfloat16`,数据格式支持ND。
80-- **query\_rope**(`Tensor`):表示Query位置编码的输出Tensor,即公式中q<sup>R</sup>。shape支持3维和4维,格式为\(T, N, Dr\)和\(B, S, N, Dr\),dtype支持`bfloat16`,数据格式支持ND。79+- **query\_rope**(`Tensor`):表示Query位置编码的输出Tensor,即公式中q<sup>R</sup>。shape支持3维和4维,格式为\(T, N, Dr\)和\(B, S, N, Dr\),dtype支持`bfloat16`,数据格式支持ND。
81-- **kv\_cache\_out**(`Tensor`):表示Key输出到`kv_cache`中的Tensor(本质in-place更新),即公式中k<sup>C</sup>。shape支持4维,格式为\(BlockNum, BlockSize, Nkv, Hckv\),dtype支持`bfloat16`和`int8`,数据格式支持ND。80+- **kv\_cache\_out**(`Tensor`):表示Key输出到`kv_cache`中的Tensor(本质in-place更新),即公式中k<sup>C</sup>。shape支持4维,格式为\(BlockNum, BlockSize, Nkv, Hckv\),dtype支持`bfloat16`和`int8`,数据格式支持ND。
82-- **kr\_cache\_out**(`Tensor`):表示Key的位置编码输出到`kr_cache`中的Tensor(本质in-place更新),即公式中k<sup>R</sup>。shape支持4维,格式为\(BlockNum, BlockSize, Nkv, Dr\),dtype支持`bfloat16`和`int8`,数据格式支持ND。81+- **kr\_cache\_out**(`Tensor`):表示Key的位置编码输出到`kr_cache`中的Tensor(本质in-place更新),即公式中k<sup>R</sup>。shape支持4维,格式为\(BlockNum, BlockSize, Nkv, Dr\),dtype支持`bfloat16`和`int8`,数据格式支持ND。
83 82 
84## 约束说明83## 约束说明
85 84 
86-- 该接口支持推理场景下使用。85+- 该接口支持推理场景下使用。
87-- 该接口支持图模式。86+- 该接口支持图模式。
88-- 接口参数中shape格式字段含义:87+- 接口参数中shape格式字段含义:
89- - B:Batch表示输入样本批量大小,取值范围为0\~65536。88+ - B:Batch表示输入样本批量大小,取值范围为0\~65536。
90- - S:Seq-Length表示输入样本序列长度,取值范围为0\~16。89+ - S:Seq-Length表示输入样本序列长度,取值范围为0\~16。
91- - He:Head-Size表示隐藏层的大小,取值为7168。90+ - He:Head-Size表示隐藏层的大小,取值为7168。
92 91 
93- - Hcq:q低秩矩阵维度,取值为1536。92+ - Hcq:q低秩矩阵维度,取值为1536。
94- - N:Head-Num表示多头数,取值范围为1、2、4、8、16、32、64、128。93+ - N:Head-Num表示多头数,取值范围为1、2、4、8、16、32、64、128。
95 94 
96- - Hckv:kv低秩矩阵维度,取值为512。95+ - Hckv:kv低秩矩阵维度,取值为512。
97- - D:qk不含位置编码维度,取值为128。96+ - D:qk不含位置编码维度,取值为128。
98- - Dr:qk位置编码维度,取值为64。97+ - Dr:qk位置编码维度,取值为64。
99- - Nkv:kv的head数,取值为1。98+ - Nkv:kv的head数,取值为1。
100- - BlockNum:PagedAttention场景下的块数,取值为计算B\*Skv/BlockSize的值后再向上取整,其中Skv表示kv的序列长度,该值允许取0。99+ - BlockNum:PagedAttention场景下的块数,取值为计算B\*Skv/BlockSize的值后再向上取整,其中Skv表示kv的序列长度,该值允许取0。
101- - BlockSize:PagedAttention场景下的块大小,取值范围为16、128。100+ - BlockSize:PagedAttention场景下的块大小,取值范围为16、128。
102- - T:BS合轴后的大小,取值范围:0\~1048576。注:若采用BS合轴,此时token\_x、rope\_sin、rope\_cos均为2维,cache\_index为1维,query、query\_rope为3维。101+ - T:BS合轴后的大小,取值范围:0\~1048576。注:若采用BS合轴,此时token\_x、rope\_sin、rope\_cos均为2维,cache\_index为1维,query、query\_rope为3维。
103-- shape约束102+- shape约束
104- - B、S、T、Skv值允许一个或多个取0,即Shape与B、S、T、Skv值相关的入参允许传入空Tensor,其余入参不支持传入空Tensor。103+ - B、S、T、Skv值允许一个或多个取0,即Shape与B、S、T、Skv值相关的入参允许传入空Tensor,其余入参不支持传入空Tensor。
105- - 如果B、S、T取值为0,则query、query_rope输出空Tensor,kv_cache、kr_cache、kv_cache_out、kr_cache_out不更新。104+ - 如果B、S、T取值为0,则query、query_rope输出空Tensor,kv_cache、kr_cache、kv_cache_out、kr_cache_out不更新。
106- - 如果Skv取值为0,则query、query_rope正常计算,kv_cache、kr_cache、kv_cache_out、kr_cache_out不更新,即输出空Tensor。105+ - 如果Skv取值为0,则query、query_rope正常计算,kv_cache、kr_cache、kv_cache_out、kr_cache_out不更新,即输出空Tensor。
107 106 
108## 调用示例107## 调用示例
109 108 
110-- 单算子模式调用109+- 单算子模式调用
111 110 
112 ```python111 ```python
113 import torch112 import torch
@@ -158,7 +157,7 @@ torch_npu.npu_mla_prolog(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv
158 device='npu:0', dtype=torch.bfloat16)157 device='npu:0', dtype=torch.bfloat16)
159 ```158 ```
160 159 
161-- 图模式调用160+- 图模式调用
162 161 
163 ```python162 ```python
164 # 入图方式163 # 入图方式
@@ -250,4 +249,3 @@ torch_npu.npu_mla_prolog(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv
250 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],249 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],
251 device='npu:0', dtype=torch.bfloat16) 250 device='npu:0', dtype=torch.bfloat16)
252 ```251 ```
253- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_mla_prolog_v2.md+44-45
@@ -12,41 +12,41 @@
12 12 
13## 功能说明13## 功能说明
14 14 
15-- API功能:推理场景下,Multi-Head Latent Attention(MLA)前处理的计算。主要计算过程分为五路;15+- API功能:推理场景下,Multi-Head Latent Attention(MLA)前处理的计算。主要计算过程分为五路;
16- - 首先对输入x乘以W<sup>DQ</sup>进行下采样和RmsNorm后分成两路,第一路乘以W<sup>UQ</sup>和W<sup>UK</sup>经过两次上采样后得到q<sup>N</sup>;16+ - 首先对输入x乘以W<sup>DQ</sup>进行下采样和RmsNorm后分成两路,第一路乘以W<sup>UQ</sup>和W<sup>UK</sup>经过两次上采样后得到q<sup>N</sup>;
17- - 第二路乘以W<sup>QR</sup>后经过旋转位置编码(ROPE)得到q<sup>R</sup>。17+ - 第二路乘以W<sup>QR</sup>后经过旋转位置编码(ROPE)得到q<sup>R</sup>。
18- - 第三路是输入x乘以W<sup>DKV</sup>进行下采样和RmsNorm后传入Cache中得到k<sup>C</sup>。18+ - 第三路是输入x乘以W<sup>DKV</sup>进行下采样和RmsNorm后传入Cache中得到k<sup>C</sup>。
19- - 第四路是输入x乘以W<sup>KR</sup>后经过旋转位置编码后传入另一个Cache中得到k<sup>R</sup>。19+ - 第四路是输入x乘以W<sup>KR</sup>后经过旋转位置编码后传入另一个Cache中得到k<sup>R</sup>。
20- - 第五路是输出q<sup>N</sup>经过DynamicQuant后得到的量化参数。20+ - 第五路是输出q<sup>N</sup>经过DynamicQuant后得到的量化参数。
21 21 
22-- 计算公式:22+- 计算公式:
23- - RmsNorm公式23+ - RmsNorm公式
24 $$RmsNorm(x) = \gamma \cdot \frac{x} {\sqrt{\frac{1}{N} \sum_{i=1}^{N} x_i^2 + \epsilon}}$$24 $$RmsNorm(x) = \gamma \cdot \frac{x} {\sqrt{\frac{1}{N} \sum_{i=1}^{N} x_i^2 + \epsilon}}$$
25 25 
26- - Query计算公式26+ - Query计算公式
27 $$c^Q = RmsNorm(x \cdot W^{DQ})$$27 $$c^Q = RmsNorm(x \cdot W^{DQ})$$
28 $$q^C = c^Q \cdot W^{UQ}$$28 $$q^C = c^Q \cdot W^{UQ}$$
29 $$q^N = q^C \cdot W^{UK}$$29 $$q^N = q^C \cdot W^{UK}$$
30 30 
31- - Query ROPE旋转位置编码31+ - Query ROPE旋转位置编码
32 32 
33 $$q^R = ROPE(c^Q \cdot W^{QR})$$33 $$q^R = ROPE(c^Q \cdot W^{QR})$$
34 34 
35- - Key计算公式35+ - Key计算公式
36 36 
37 $$k^C = Cache(RmsNorm(x \cdot W^{DKV}))$$37 $$k^C = Cache(RmsNorm(x \cdot W^{DKV}))$$
38 38 
39- - Key ROPE旋转位置编码39+ - Key ROPE旋转位置编码
40 40 
41 $$k^R = Cache(ROPE(x \cdot W^{KR}))$$41 $$k^R = Cache(ROPE(x \cdot W^{KR}))$$
42 42 
43- - 反量化缩放因子(Dequant Scale Query Nope)计算公式43+ - 反量化缩放因子(Dequant Scale Query Nope)计算公式
44 $$\text{dequantScaleQNope} = \frac{\text{RowMax}(\text{abs}(q^N))}{127}$$44 $$\text{dequantScaleQNope} = \frac{\text{RowMax}(\text{abs}(q^N))}{127}$$
45 $$q^N = \text{round}(\frac{q^N}{\text{dequantScaleQNope}})$$45 $$q^N = \text{round}(\frac{q^N}{\text{dequantScaleQNope}})$$
46 46 
47## 函数原型47## 函数原型
48 48 
49-```49+```python
50torch_npu.npu_mla_prolog_v2(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv_kr, rmsnorm_gamma_cq, rmsnorm_gamma_ckv, rope_sin, rope_cos, cache_index, kv_cache, kr_cache, *, dequant_scale_x=None, dequant_scale_w_dq=None, dequant_scale_w_uq_qr=None, dequant_scale_w_dkv_kr=None, quant_scale_ckv=None, quant_scale_ckr=None, smooth_scales_cq=None, rmsnorm_epsilon_cq=1e-05, rmsnorm_epsilon_ckv=1e-05, cache_mode="PA_BSND") -> (Tensor, Tensor, Tensor, Tensor, Tensor)50torch_npu.npu_mla_prolog_v2(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv_kr, rmsnorm_gamma_cq, rmsnorm_gamma_ckv, rope_sin, rope_cos, cache_index, kv_cache, kr_cache, *, dequant_scale_x=None, dequant_scale_w_dq=None, dequant_scale_w_uq_qr=None, dequant_scale_w_dkv_kr=None, quant_scale_ckv=None, quant_scale_ckr=None, smooth_scales_cq=None, rmsnorm_epsilon_cq=1e-05, rmsnorm_epsilon_ckv=1e-05, cache_mode="PA_BSND") -> (Tensor, Tensor, Tensor, Tensor, Tensor)
51```51```
52 52 
@@ -78,40 +78,40 @@ torch_npu.npu_mla_prolog_v2(token_x, weight_dq, weight_uq_qr, weight_uk, weight_
78 78 
79## 返回值说明79## 返回值说明
80 80 
81-- **query**(`Tensor`):表示Query的输出Tensor,即公式中q<sup>N</sup>。shape支持3维和4维,格式为\(T, N, Hckv\)和\(B, S, N, Hckv\),dtype支持`bfloat16`和`int8`,数据格式支持ND。81+- **query**(`Tensor`):表示Query的输出Tensor,即公式中q<sup>N</sup>。shape支持3维和4维,格式为\(T, N, Hckv\)和\(B, S, N, Hckv\),dtype支持`bfloat16`和`int8`,数据格式支持ND。
82-- **query\_rope**(`Tensor`):表示Query位置编码的输出Tensor,即公式中q<sup>R</sup>。shape支持3维和4维,格式为\(T, N, Dr\)和\(B, S, N, Dr\),dtype支持`bfloat16`,数据格式支持ND。82+- **query\_rope**(`Tensor`):表示Query位置编码的输出Tensor,即公式中q<sup>R</sup>。shape支持3维和4维,格式为\(T, N, Dr\)和\(B, S, N, Dr\),dtype支持`bfloat16`,数据格式支持ND。
83-- **kv\_cache\_out**(`Tensor`):表示Key输出到`kv_cache`中的Tensor(本质in-place更新),即公式中k<sup>C</sup>。shape支持4维,格式为\(BlockNum, BlockSize, Nkv, Hckv\),dtype支持`bfloat16`和`int8`,数据格式支持ND。83+- **kv\_cache\_out**(`Tensor`):表示Key输出到`kv_cache`中的Tensor(本质in-place更新),即公式中k<sup>C</sup>。shape支持4维,格式为\(BlockNum, BlockSize, Nkv, Hckv\),dtype支持`bfloat16`和`int8`,数据格式支持ND。
84-- **kr\_cache\_out**(`Tensor`):表示Key的位置编码输出到`kr_cache`中的Tensor(本质in-place更新),即公式中k<sup>R</sup>。shape支持4维,格式为\(BlockNum, BlockSize, Nkv, Dr\),dtype支持`bfloat16`和`int8`,数据格式支持ND。84+- **kr\_cache\_out**(`Tensor`):表示Key的位置编码输出到`kr_cache`中的Tensor(本质in-place更新),即公式中k<sup>R</sup>。shape支持4维,格式为\(BlockNum, BlockSize, Nkv, Dr\),dtype支持`bfloat16`和`int8`,数据格式支持ND。
85-- **dequant\_scale\_q\_nope**(`Tensor`):表示Query的输出Tensor的反量化参数。其shape支持1维和3维,全量化kv\_cache量化场景下,其shape为\(T, N, 1\)和\(B\*S, N, 1\);其他场景下,其shape为\(1\),dtype支持`float`,数据格式支持ND。85+- **dequant\_scale\_q\_nope**(`Tensor`):表示Query的输出Tensor的反量化参数。其shape支持1维和3维,全量化kv\_cache量化场景下,其shape为\(T, N, 1\)和\(B\*S, N, 1\);其他场景下,其shape为\(1\),dtype支持`float`,数据格式支持ND。
86 86 
87## 约束说明87## 约束说明
88 88 
89-- 该接口支持推理场景下使用。89+- 该接口支持推理场景下使用。
90-- 该接口支持图模式。90+- 该接口支持图模式。
91-- 接口参数中shape格式字段含义:91+- 接口参数中shape格式字段含义:
92- - B:Batch表示输入样本批量大小,取值范围为0\~65536。92+ - B:Batch表示输入样本批量大小,取值范围为0\~65536。
93- - S:Seq-Length表示输入样本序列长度,取值范围为0\~16。93+ - S:Seq-Length表示输入样本序列长度,取值范围为0\~16。
94- - He:Head-Size表示隐藏层的大小,取值为7168。94+ - He:Head-Size表示隐藏层的大小,取值为7168。
95 95 
96- - Hcq:q低秩矩阵维度,取值为1536。96+ - Hcq:q低秩矩阵维度,取值为1536。
97- - N:Head-Num表示多头数,取值范围为1、2、4、8、16、32、64、128。97+ - N:Head-Num表示多头数,取值范围为1、2、4、8、16、32、64、128。
98 98 
99- - Hckv:kv低秩矩阵维度,取值为512。99+ - Hckv:kv低秩矩阵维度,取值为512。
100- - D:qk不含位置编码维度,取值为128。100+ - D:qk不含位置编码维度,取值为128。
101- - Dr:qk位置编码维度,取值为64。101+ - Dr:qk位置编码维度,取值为64。
102- - Nkv:kv的head数,取值为1。102+ - Nkv:kv的head数,取值为1。
103- - BlockNum:PagedAttention场景下的块数,取值为计算B\*Skv/BlockSize的值后再向上取整,其中Skv表示kv的序列长度,该值允许取0。103+ - BlockNum:PagedAttention场景下的块数,取值为计算B\*Skv/BlockSize的值后再向上取整,其中Skv表示kv的序列长度,该值允许取0。
104- - BlockSize:PagedAttention场景下的块大小,取值范围为16、128。104+ - BlockSize:PagedAttention场景下的块大小,取值范围为16、128。
105- - T:BS合轴后的大小,取值范围:0\~1048576。105+ - T:BS合轴后的大小,取值范围:0\~1048576。
106 106 
107-- shape约束:107+- shape约束:
108- - 若`token_x`的维度采用BS合轴,即\(T, He\),则`rope_sin`和rope\_cos的shape为\(T, Dr\),cache\_index的shape为\(T,\),dequant\_scale\_x的shape为\(T, 1\),query的shape为\(T, N, Hckv\),query\_rope的shape为\(T, N, Dr\)。全量化kv\_cache量化场景下,dequant\_scale\_q\_nope的shape为\(T, N, 1\),其他场景下dequant\_scale\_q\_nope的shape为\(1\)。108+ - 若`token_x`的维度采用BS合轴,即\(T, He\),则`rope_sin`和rope\_cos的shape为\(T, Dr\),cache\_index的shape为\(T,\),dequant\_scale\_x的shape为\(T, 1\),query的shape为\(T, N, Hckv\),query\_rope的shape为\(T, N, Dr\)。全量化kv\_cache量化场景下,dequant\_scale\_q\_nope的shape为\(T, N, 1\),其他场景下dequant\_scale\_q\_nope的shape为\(1\)。
109- - 若`token_x`的维度不采用BS合轴,即\(B, S, He\),则`rope_sin`和rope\_cos的shape为\(B, S, Dr\),cache\_index的shape为\(B, S\),dequant\_scale\_x的shape为\(B\*S, 1\),query的shape为\(B, S, N, Hckv\),query\_rope的shape为\(B, S, N, Dr\)。全量化kv\_cache量化场景下,dequant\_scale\_q\_nope的shape为\(B\*S, N, 1\),其他场景下dequant\_scale\_q\_nope的shape为\(1\)。109+ - 若`token_x`的维度不采用BS合轴,即\(B, S, He\),则`rope_sin`和rope\_cos的shape为\(B, S, Dr\),cache\_index的shape为\(B, S\),dequant\_scale\_x的shape为\(B\*S, 1\),query的shape为\(B, S, N, Hckv\),query\_rope的shape为\(B, S, N, Dr\)。全量化kv\_cache量化场景下,dequant\_scale\_q\_nope的shape为\(B\*S, N, 1\),其他场景下dequant\_scale\_q\_nope的shape为\(1\)。
110- - B、S、T、Skv值允许一个或多个取0,即Shape与B、S、T、Skv值相关的入参允许传入空Tensor,其余入参不支持传入空Tensor。110+ - B、S、T、Skv值允许一个或多个取0,即Shape与B、S、T、Skv值相关的入参允许传入空Tensor,其余入参不支持传入空Tensor。
111- - 如果B、S、T取值为0,则query、query_rope、dequant_scale_q_nope输出空Tensor,kv_cache、kr_cache、kv_cache_out、kr_cache_out不更新。111+ - 如果B、S、T取值为0,则query、query_rope、dequant_scale_q_nope输出空Tensor,kv_cache、kr_cache、kv_cache_out、kr_cache_out不更新。
112- - 如果Skv取值为0,则query、query_rope、dequant_scale_q_nope正常计算,kv_cache、kr_cache、kv_cache_out、kr_cache_out不更新,即输出空Tensor。112+ - 如果Skv取值为0,则query、query_rope、dequant_scale_q_nope正常计算,kv_cache、kr_cache、kv_cache_out、kr_cache_out不更新,即输出空Tensor。
113 113 
114-- 本算子支持以下场景:114+- 本算子支持以下场景:
115 <a name="zh-cn_topic_0000002313328922_table664817810310"></a>115 <a name="zh-cn_topic_0000002313328922_table664817810310"></a>
116 <table><thead align="left"><tr id="zh-cn_topic_0000002313328922_row9649788313"><th class="cellrowborder" colspan="2" valign="top" id="mcps1.1.4.1.1"><p id="zh-cn_topic_0000002313328922_p14649381739"><a name="zh-cn_topic_0000002313328922_p14649381739"></a><a name="zh-cn_topic_0000002313328922_p14649381739"></a>场景</p>116 <table><thead align="left"><tr id="zh-cn_topic_0000002313328922_row9649788313"><th class="cellrowborder" colspan="2" valign="top" id="mcps1.1.4.1.1"><p id="zh-cn_topic_0000002313328922_p14649381739"><a name="zh-cn_topic_0000002313328922_p14649381739"></a><a name="zh-cn_topic_0000002313328922_p14649381739"></a>场景</p>
117 </th>117 </th>
@@ -155,7 +155,7 @@ torch_npu.npu_mla_prolog_v2(token_x, weight_dq, weight_uq_qr, weight_uk, weight_
155 </tbody>155 </tbody>
156 </table>156 </table>
157 157 
158-- 在不同量化场景下,参数的dtype和shape组合需满足如下条件:158+- 在不同量化场景下,参数的dtype和shape组合需满足如下条件:
159 159 
160 <a name="zh-cn_topic_0000002313328922_table1311951423117"></a>160 <a name="zh-cn_topic_0000002313328922_table1311951423117"></a>
161 <table><tbody><tr id="zh-cn_topic_0000002313328922_row510181463115"><td class="cellrowborder" rowspan="3" valign="top"><p id="zh-cn_topic_0000002313328922_p21013144313"><a name="zh-cn_topic_0000002313328922_p21013144313"></a><a name="zh-cn_topic_0000002313328922_p21013144313"></a><strong id="zh-cn_topic_0000002313328922_b1423515521358"><a name="zh-cn_topic_0000002313328922_b1423515521358"></a><a name="zh-cn_topic_0000002313328922_b1423515521358"></a>参数名</strong></p>161 <table><tbody><tr id="zh-cn_topic_0000002313328922_row510181463115"><td class="cellrowborder" rowspan="3" valign="top"><p id="zh-cn_topic_0000002313328922_p21013144313"><a name="zh-cn_topic_0000002313328922_p21013144313"></a><a name="zh-cn_topic_0000002313328922_p21013144313"></a><strong id="zh-cn_topic_0000002313328922_b1423515521358"><a name="zh-cn_topic_0000002313328922_b1423515521358"></a><a name="zh-cn_topic_0000002313328922_b1423515521358"></a>参数名</strong></p>
@@ -754,7 +754,7 @@ torch_npu.npu_mla_prolog_v2(token_x, weight_dq, weight_uq_qr, weight_uk, weight_
754 754 
755## 调用示例<a name="zh-cn_topic_0000002313328922_section983519211229"></a>755## 调用示例<a name="zh-cn_topic_0000002313328922_section983519211229"></a>
756 756 
757-- 单算子模式调用757+- 单算子模式调用
758 758 
759 ```python759 ```python
760 import torch760 import torch
@@ -807,7 +807,7 @@ torch_npu.npu_mla_prolog_v2(token_x, weight_dq, weight_uq_qr, weight_uk, weight_
807 device='npu:0', dtype=torch.bfloat16)807 device='npu:0', dtype=torch.bfloat16)
808 ```808 ```
809 809 
810-- 图模式调用810+- 图模式调用
811 811 
812 ```python812 ```python
813 # 入图方式813 # 入图方式
@@ -898,4 +898,3 @@ torch_npu.npu_mla_prolog_v2(token_x, weight_dq, weight_uq_qr, weight_uk, weight_
898 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],898 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],
899 device='npu:0', dtype=torch.bfloat16) 899 device='npu:0', dtype=torch.bfloat16)
900 ```900 ```
901- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_mla_prolog_v3.md+61-60
@@ -8,22 +8,22 @@
8 8 
9## 功能说明9## 功能说明
10 10 
11-- API功能:推理场景下Multi-Head Latent Attention前处理的计算操作。该算子实现四条并行的计算路径:11+- API功能:推理场景下Multi-Head Latent Attention前处理的计算操作。该算子实现四条并行的计算路径:
12 1. 标准Query路径:输入$x$ → $W^{DQ}$下采样 → RmsNorm → $W^{UQ}$上采样 → $W^{UK}$上采样 → $q^N$12 1. 标准Query路径:输入$x$ → $W^{DQ}$下采样 → RmsNorm → $W^{UQ}$上采样 → $W^{UK}$上采样 → $q^N$
13 2. 位置编码Query路径:输入$x$ → $W^{DQ}$下采样 → RmsNorm → $W^{QR}$ → ROPE旋转位置编码 → $q^R$13 2. 位置编码Query路径:输入$x$ → $W^{DQ}$下采样 → RmsNorm → $W^{QR}$ → ROPE旋转位置编码 → $q^R$
14 3. 标准Key路径:输入$x$ → $W^{DKV}$下采样 → RmsNorm → Cache存储 → $k^C$14 3. 标准Key路径:输入$x$ → $W^{DKV}$下采样 → RmsNorm → Cache存储 → $k^C$
15 4. 位置编码Key路径:输入$x$ → $W^{KR}$ → ROPE旋转位置编码 → Cache存储 → $k^R$15 4. 位置编码Key路径:输入$x$ → $W^{KR}$ → ROPE旋转位置编码 → Cache存储 → $k^R$
16 16 
17-- 相比torch_npu.npu_mla_prolog_v2的主要差异如下:17+- 相比torch_npu.npu_mla_prolog_v2的主要差异如下:
18- - 新增输出`query_norm`和`dequant_scale_q_norm`,用于支持DeepSeekV3.2网络。18+ - 新增输出`query_norm`和`dequant_scale_q_norm`,用于支持DeepSeekV3.2网络。
19- - 新增`kv_cache`的per-tile量化模式。19+ - 新增`kv_cache`的per-tile量化模式。
20- - 新增query与key的尺度矫正因子,分别对应qc_qr_scale($\alpha_q$)与kc_scale($\alpha_{kv}$)。20+ - 新增query与key的尺度矫正因子,分别对应qc_qr_scale($\alpha_q$)与kc_scale($\alpha_{kv}$)。
21- - 新增`cache_mode`对"PA_BLK_BSND"、"PA_BLK_NZ"、"BSND"和"TND"格式的支持。21+ - 新增`cache_mode`对"PA_BLK_BSND"、"PA_BLK_NZ"、"BSND"和"TND"格式的支持。
22- - 新增可选参数`weight_quant_mode`、`kv_cache_quant_mode`、`query_quant_mode`、`ckvkr_repo_mode`、`quant_scale_repo_mode`,用于配置量化场景。22+ - 新增可选参数`weight_quant_mode`、`kv_cache_quant_mode`、`query_quant_mode`、`ckvkr_repo_mode`、`quant_scale_repo_mode`,用于配置量化场景。
23- - 调整`cache_index`为可选参数。23+ - 调整`cache_index`为可选参数。
24 24 
25-- 计算公式:25+- 计算公式:
26- - RmsNorm公式26+ - RmsNorm公式
27 $$27 $$
28 \text{RmsNorm}(x) = \gamma \cdot \frac{x_i}{\text{RMS}(x)}28 \text{RmsNorm}(x) = \gamma \cdot \frac{x_i}{\text{RMS}(x)}
29 $$29 $$
@@ -32,7 +32,7 @@
32 \text{RMS}(x) = \sqrt{\frac{1}{N} \sum_{i=1}^{N} x_i^2 + \epsilon}32 \text{RMS}(x) = \sqrt{\frac{1}{N} \sum_{i=1}^{N} x_i^2 + \epsilon}
33 $$33 $$
34 34 
35- - 路径1:标准Query计算35+ - 路径1:标准Query计算
36 36 
37 包括下采样、RmsNorm和两次上采样:37 包括下采样、RmsNorm和两次上采样:
38 38 
@@ -48,7 +48,7 @@
48 q^N = q^C \cdot W^{UK}48 q^N = q^C \cdot W^{UK}
49 $$49 $$
50 50 
51- - 路径2:位置编码Query计算51+ - 路径2:位置编码Query计算
52 52 
53 对Query进行ROPE旋转位置编码:53 对Query进行ROPE旋转位置编码:
54 54 
@@ -56,7 +56,7 @@
56 q^R = ROPE(c^Q \cdot W^{QR})56 q^R = ROPE(c^Q \cdot W^{QR})
57 $$57 $$
58 58 
59- - 路径3:标准Key计算59+ - 路径3:标准Key计算
60 60 
61 包括下采样、RmsNorm,将计算结果存入Cache:61 包括下采样、RmsNorm,将计算结果存入Cache:
62 62 
@@ -68,7 +68,7 @@
68 k^C = Cache(c^{KV})68 k^C = Cache(c^{KV})
69 $$69 $$
70 70 
71- - 路径4:位置编码Key计算71+ - 路径4:位置编码Key计算
72 72 
73 对Key进行ROPE旋转位置编码,并将结果存入Cache:73 对Key进行ROPE旋转位置编码,并将结果存入Cache:
74 74 
@@ -76,9 +76,9 @@
76 k^R = Cache(ROPE(x \cdot W^{KR}))76 k^R = Cache(ROPE(x \cdot W^{KR}))
77 $$77 $$
78 78 
79- 
80## 函数原型79## 函数原型
81-```80+ 
81+```python
82torch_npu.npu_mla_prolog_v3(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv_kr, rmsnorm_gamma_cq, rmsnorm_gamma_ckv, rope_sin, rope_cos, kv_cache, kr_cache, cache_index=None, dequant_scale_x=None, dequant_scale_w_dq=None, dequant_scale_w_uq_qr=None, dequant_scale_w_dkv_kr=None, quant_scale_ckv=None, quant_scale_ckr=None, smooth_scales_cq=None, actual_seq_len=None, k_nope_clip_alpha=None, rmsnorm_epsilon_cq=1e-05, rmsnorm_epsilon_ckv=1e-05, cache_mode='PA_BSND', query_norm_flag=False, weight_quant_mode=0, kv_cache_quant_mode=0, query_quant_mode=0, ckvkr_repo_mode=0, quant_scale_repo_mode=0, tile_size=128, qc_qr_scale=1.0, kc_scale=1.0) -> (Tensor, Tensor, Tensor, Tensor, Tensor)82torch_npu.npu_mla_prolog_v3(token_x, weight_dq, weight_uq_qr, weight_uk, weight_dkv_kr, rmsnorm_gamma_cq, rmsnorm_gamma_ckv, rope_sin, rope_cos, kv_cache, kr_cache, cache_index=None, dequant_scale_x=None, dequant_scale_w_dq=None, dequant_scale_w_uq_qr=None, dequant_scale_w_dkv_kr=None, quant_scale_ckv=None, quant_scale_ckr=None, smooth_scales_cq=None, actual_seq_len=None, k_nope_clip_alpha=None, rmsnorm_epsilon_cq=1e-05, rmsnorm_epsilon_ckv=1e-05, cache_mode='PA_BSND', query_norm_flag=False, weight_quant_mode=0, kv_cache_quant_mode=0, query_quant_mode=0, ckvkr_repo_mode=0, quant_scale_repo_mode=0, tile_size=128, qc_qr_scale=1.0, kc_scale=1.0) -> (Tensor, Tensor, Tensor, Tensor, Tensor)
83```83```
84 84 
@@ -87,91 +87,94 @@ torch_npu.npu_mla_prolog_v3(token_x, weight_dq, weight_uq_qr, weight_uk, weight_
87> [!NOTE]87> [!NOTE]
88> B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、He(Head Size)表示隐藏层大小、N(Head Num)表示多头数、Hcq表示q低秩矩阵维度、Hckv表示kv低秩矩阵维度、Dtile表示kv_cache的D轴维度、D表示qk不含位置编码维度、Dr表示qk位置编码维度、Nkv表示kv的head数、BlockNum表示PagedAttention场景下的块数、BlockSize表示PagedAttention场景下的块大小、T表示BS合轴后的大小。88> B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、He(Head Size)表示隐藏层大小、N(Head Num)表示多头数、Hcq表示q低秩矩阵维度、Hckv表示kv低秩矩阵维度、Dtile表示kv_cache的D轴维度、D表示qk不含位置编码维度、Dr表示qk位置编码维度、Nkv表示kv的head数、BlockNum表示PagedAttention场景下的块数、BlockSize表示PagedAttention场景下的块大小、T表示BS合轴后的大小。
89 89 
90-- **token_x**(`Tensor`):必选参数,公式中用于计算Query和Key的输入tensor。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`、`int8`。BS合轴时,shape为[T, He];BS非合轴时,shape为[B, S, He]。90+- **token_x**(`Tensor`):必选参数,公式中用于计算Query和Key的输入tensor。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`、`int8`。BS合轴时,shape为[T, He];BS非合轴时,shape为[B, S, He]。
91 91 
92-- **weight_dq**(`Tensor`):必选参数,公式中用于计算Query的下采样权重矩阵$W^{DQ}$。不支持非连续,数据格式支持FRACTAL_NZ,数据类型支持`bfloat16`、`int8`,shape为[He, Hcq]。92+- **weight_dq**(`Tensor`):必选参数,公式中用于计算Query的下采样权重矩阵$W^{DQ}$。不支持非连续,数据格式支持FRACTAL_NZ,数据类型支持`bfloat16`、`int8`,shape为[He, Hcq]。
93 93 
94-- **weight_uq_qr**(`Tensor`):必选参数,公式中用于计算Query的上采样权重矩阵$W^{UQ}$和位置编码权重矩阵$W^{QR}$。不支持非连续,数据格式支持FRACTAL_NZ,数据类型支持`bfloat16`、`int8`,shape为[Hcq, N*(D+Dr)]。94+- **weight_uq_qr**(`Tensor`):必选参数,公式中用于计算Query的上采样权重矩阵$W^{UQ}$和位置编码权重矩阵$W^{QR}$。不支持非连续,数据格式支持FRACTAL_NZ,数据类型支持`bfloat16`、`int8`,shape为[Hcq, N*(D+Dr)]。
95 95 
96-- **weight_uk**(`Tensor`):必选参数,公式中用于计算Key的上采样权重$W^{UK}$。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`,shape为[N, D, Hckv]。96+- **weight_uk**(`Tensor`):必选参数,公式中用于计算Key的上采样权重$W^{UK}$。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`,shape为[N, D, Hckv]。
97 97 
98-- **weight_dkv_kr**(`Tensor`):必选参数,公式中用于计算Key的下采样权重矩阵$W^{DKV}$和位置编码权重矩阵$W^{KR}$。不支持非连续,数据格式支持FRACTAL_NZ,数据类型支持`bfloat16`、`int8`,shape为[He, Hckv+Dr]。98+- **weight_dkv_kr**(`Tensor`):必选参数,公式中用于计算Key的下采样权重矩阵$W^{DKV}$和位置编码权重矩阵$W^{KR}$。不支持非连续,数据格式支持FRACTAL_NZ,数据类型支持`bfloat16`、`int8`,shape为[He, Hckv+Dr]。
99 99 
100-- **rmsnorm_gamma_cq**(`Tensor`):必选参数,计算$c^Q$的RmsNorm公式中的$\gamma$参数。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`,shape为[Hcq]。100+- **rmsnorm_gamma_cq**(`Tensor`):必选参数,计算$c^Q$的RmsNorm公式中的$\gamma$参数。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`,shape为[Hcq]。
101 101 
102-- **rmsnorm_gamma_ckv**(`Tensor`):必选参数,计算$c^{KV}$的RmsNorm公式中的$\gamma$参数。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`,shape为[Hckv]。102+- **rmsnorm_gamma_ckv**(`Tensor`):必选参数,计算$c^{KV}$的RmsNorm公式中的$\gamma$参数。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`,shape为[Hckv]。
103 103 
104-- **rope_sin**(`Tensor`):必选参数,用于计算旋转位置编码的正弦参数矩阵。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`。BS合轴时,shape为[T, Dr];BS非合轴时,shape为[B, S, Dr]。支持B=0,S=0,T=0的空Tensor。104+- **rope_sin**(`Tensor`):必选参数,用于计算旋转位置编码的正弦参数矩阵。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`。BS合轴时,shape为[T, Dr];BS非合轴时,shape为[B, S, Dr]。支持B=0,S=0,T=0的空Tensor。
105 105 
106-- **rope_cos**(`Tensor`):必选参数,用于计算旋转位置编码的余弦参数矩阵。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`;BS合轴时,shape为[T, Dr];BS非合轴时,shape为[B, S, Dr]。支持B=0,S=0,T=0的空Tensor。106+- **rope_cos**(`Tensor`):必选参数,用于计算旋转位置编码的余弦参数矩阵。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`;BS合轴时,shape为[T, Dr];BS非合轴时,shape为[B, S, Dr]。支持B=0,S=0,T=0的空Tensor。
107 107 
108-- **kv_cache**(`Tensor`):必选参数,表示cache的索引,计算结果原地更新(对应公式中的$k^C$)。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`、`int8`,当cache_mode为"PA_BSND"、"PA_NZ"、"PA_BLK_BSND"、"PA_BLK_NZ"时shape为[BlockNum, BlockSize, Nkv, Dtile],支持B=0,Skv=0的空Tensor;当cache_mode为"BSND"时shape为[B, S, Nkv, Dtile],不支持空Tensor;当cache_mode为"TND"时shape为[T, Nkv, Dtile],不支持空Tensor;Nkv与N关联,N是超参,故不支持Nkv=0。108+- **kv_cache**(`Tensor`):必选参数,表示cache的索引,计算结果原地更新(对应公式中的$k^C$)。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`、`int8`,当cache_mode为"PA_BSND"、"PA_NZ"、"PA_BLK_BSND"、"PA_BLK_NZ"时shape为[BlockNum, BlockSize, Nkv, Dtile],支持B=0,Skv=0的空Tensor;当cache_mode为"BSND"时shape为[B, S, Nkv, Dtile],不支持空Tensor;当cache_mode为"TND"时shape为[T, Nkv, Dtile],不支持空Tensor;Nkv与N关联,N是超参,故不支持Nkv=0。
109 109 
110-- **kr_cache**(`Tensor`):必选参数,用于key位置编码的cache,计算结果原地更新(对应公式中的$k^R$)。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`、`int8`,当cache_mode为"PA_BSND"、"PA_NZ"、"PA_BLK_BSND"、"PA_BLK_NZ"时shape为[BlockNum, BlockSize, Nkv, Dr],支持B=0,Skv=0的空Tensor;当cache_mode为"BSND"时shape为[B, S, Nkv, Dr],不支持空Tensor;当cache_mode为"TND"时shape为[T, Nkv, Dr],不支持空Tensor;Nkv与N关联,N是超参,故不支持Nkv=0。110+- **kr_cache**(`Tensor`):必选参数,用于key位置编码的cache,计算结果原地更新(对应公式中的$k^R$)。不支持非连续,数据格式支持ND,数据类型支持`bfloat16`、`int8`,当cache_mode为"PA_BSND"、"PA_NZ"、"PA_BLK_BSND"、"PA_BLK_NZ"时shape为[BlockNum, BlockSize, Nkv, Dr],支持B=0,Skv=0的空Tensor;当cache_mode为"BSND"时shape为[B, S, Nkv, Dr],不支持空Tensor;当cache_mode为"TND"时shape为[T, Nkv, Dr],不支持空Tensor;Nkv与N关联,N是超参,故不支持Nkv=0。
111 111 
112- <strong>*</strong>:代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。112- <strong>*</strong>:代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
113 113 
114-- **cache_index**(`Tensor`):可选参数,用于存储`kv_cache`和`kr_cache`的索引。不支持非连续,数据格式支持ND,数据类型支持`int64`,当`cache_mode`为"PA_BSND"或"PA_NZ":BS合轴时shape为[T],BS非合轴时shape为[B, S],取值范围需在[0, BlockNum*BlockSize)内; 当`cache_mode`为"PA_BLK_BSND"或"PA_BLK_NZ":BS合轴时shape为[Sum(Ceil(S_i/BlockSize))](S_i表示第i个batch的序列长度),BS非合轴时shape为[B, Ceil(S/BlockSize)],取值范围需在[0, BlockNum)内;当`cache_mode`为"BSND"或"TND":`cache_index`无需传入。当前不会对传入值的合法性进行校验,需用户自行保证。114+- **cache_index**(`Tensor`):可选参数,用于存储`kv_cache`和`kr_cache`的索引。不支持非连续,数据格式支持ND,数据类型支持`int64`,当`cache_mode`为"PA_BSND"或"PA_NZ":BS合轴时shape为[T],BS非合轴时shape为[B, S],取值范围需在[0, BlockNum*BlockSize)内; 当`cache_mode`为"PA_BLK_BSND"或"PA_BLK_NZ":BS合轴时shape为[Sum(Ceil(S_i/BlockSize))](S_i表示第i个batch的序列长度),BS非合轴时shape为[B, Ceil(S/BlockSize)],取值范围需在[0, BlockNum)内;当`cache_mode`为"BSND"或"TND":`cache_index`无需传入。当前不会对传入值的合法性进行校验,需用户自行保证。
115 115 
116-- **dequant_scale_x**(`Tensor`):可选参数,token_x的反量化参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[T]或[B*S, 1],支持B=0,S=0,T=0的空Tensor。116+- **dequant_scale_x**(`Tensor`):可选参数,token_x的反量化参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[T]或[B*S, 1],支持B=0,S=0,T=0的空Tensor。
117 117 
118-- **dequant_scale_w_dq**(`Tensor`):可选参数,weight_dq的反量化参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, Hcq]。118+- **dequant_scale_w_dq**(`Tensor`):可选参数,weight_dq的反量化参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, Hcq]。
119 119 
120-- **dequant_scale_w_uq_qr**(`Tensor`):可选参数,用于MatmulQcQr矩阵乘后反量化操作的per-channel参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, N*(D+Dr)]。120+- **dequant_scale_w_uq_qr**(`Tensor`):可选参数,用于MatmulQcQr矩阵乘后反量化操作的per-channel参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, N*(D+Dr)]。
121 121 
122-- **dequant_scale_w_dkv_kr**(`Tensor`):可选参数,weight_dkv_kr的反量化参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, Hckv+Dr]。122+- **dequant_scale_w_dkv_kr**(`Tensor`):可选参数,weight_dkv_kr的反量化参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, Hckv+Dr]。
123 123 
124-- **quant_scale_ckv**(`Tensor`):可选参数,用于对kv_cache输出数据做量化操作的参数。不支持非连续,数据格式支持ND,数据类型支持`float`,*部分量化*场景时shape为[1, Hckv],*全量化*场景时shape为[1],支持非空Tensor(仅kv_cache为`int8` dtype输出场景需传)。124+- **quant_scale_ckv**(`Tensor`):可选参数,用于对kv_cache输出数据做量化操作的参数。不支持非连续,数据格式支持ND,数据类型支持`float`,*部分量化*场景时shape为[1, Hckv],*全量化*场景时shape为[1],支持非空Tensor(仅kv_cache为`int8` dtype输出场景需传)。
125 125 
126-- **quant_scale_ckr**(`Tensor`):可选参数,用于对kr_cache输出数据做量化操作的参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, Dr],支持非空Tensor(仅`int8` dtype量化输出场景需传)。126+- **quant_scale_ckr**(`Tensor`):可选参数,用于对kr_cache输出数据做量化操作的参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, Dr],支持非空Tensor(仅`int8` dtype量化输出场景需传)。
127 127 
128-- **smooth_scales_cq**(`Tensor`):可选参数,用于对RmsNorm_cq输出做动态量化操作的参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, Hcq]或[1],支持非空Tensor(仅`int8` dtype量化输出场景可选传)。128+- **smooth_scales_cq**(`Tensor`):可选参数,用于对RmsNorm_cq输出做动态量化操作的参数。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1, Hcq]或[1],支持非空Tensor(仅`int8` dtype量化输出场景可选传)。
129 129 
130-- **actual_seq_len**(`Tensor`):可选参数,表示每个batch中的序列长度,以前缀和的形式储存。不支持非连续,数据格式支持ND,数据类型支持`int32`,shape为[B],支持非空tensor(仅BS合轴且cache_mode为"PA_BLK_BSND"或"PA_BLK_NZ"时需要传入)。当前不会对传入值的合法性进行校验,需用户自行保证。130+- **actual_seq_len**(`Tensor`):可选参数,表示每个batch中的序列长度,以前缀和的形式储存。不支持非连续,数据格式支持ND,数据类型支持`int32`,shape为[B],支持非空tensor(仅BS合轴且cache_mode为"PA_BLK_BSND"或"PA_BLK_NZ"时需要传入)。当前不会对传入值的合法性进行校验,需用户自行保证。
131 131 
132-- **k_nope_clip_alpha**(`Tensor`):可选参数,表示kv_cache做clip操作时的缩放因子,当前仅在kvcache per-tile量化场景下使用。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1]。132+- **k_nope_clip_alpha**(`Tensor`):可选参数,表示kv_cache做clip操作时的缩放因子,当前仅在kvcache per-tile量化场景下使用。不支持非连续,数据格式支持ND,数据类型支持`float`,shape为[1]。
133 133 
134-- **rmsnorm_epsilon_cq**(`float`):可选参数,计算$c^Q$的RmsNorm公式中的$\epsilon$参数。默认值为1e-05。134+- **rmsnorm_epsilon_cq**(`float`):可选参数,计算$c^Q$的RmsNorm公式中的$\epsilon$参数。默认值为1e-05。
135 135 
136-- **rmsnorm_epsilon_ckv**(`float`):可选参数,计算$c^{KV}$的RmsNorm公式中的$\epsilon$参数。默认值为1e-05。136+- **rmsnorm_epsilon_ckv**(`float`):可选参数,计算$c^{KV}$的RmsNorm公式中的$\epsilon$参数。默认值为1e-05。
137 137 
138-- **cache_mode**(`str`):可选参数,表示kv_cache的模式。可选值为"PA_BSND"、"PA_NZ"、"PA_BLK_BSND"、"PA_BLK_NZ"、"TND"(对应BS合轴)和"BSND"(对应非BS合轴),默认为"PA_BSND"。138+- **cache_mode**(`str`):可选参数,表示kv_cache的模式。可选值为"PA_BSND"、"PA_NZ"、"PA_BLK_BSND"、"PA_BLK_NZ"、"TND"(对应BS合轴)和"BSND"(对应非BS合轴),默认为"PA_BSND"。
139 139 
140-- **query_norm_flag**(`bool`):可选参数,表示是否输出query_norm。仅支持bool类型,False表示不输出query_norm,True表示输出query_norm(量化场景下伴随输出dequant_scale_q_norm),默认值为False。140+- **query_norm_flag**(`bool`):可选参数,表示是否输出query_norm。仅支持bool类型,False表示不输出query_norm,True表示输出query_norm(量化场景下伴随输出dequant_scale_q_norm),默认值为False。
141 141 
142-- **weight_quant_mode**(`int`):可选参数,表示weight_dq、weight_uq_qr、weight_uk、weight_dkv_kr的量化模式。0表示非量化,1表示weight_uq_qr量化,2表示weight_dq、weight_uq_qr、weight_dkv_kr量化,默认值为0。142+- **weight_quant_mode**(`int`):可选参数,表示weight_dq、weight_uq_qr、weight_uk、weight_dkv_kr的量化模式。0表示非量化,1表示weight_uq_qr量化,2表示weight_dq、weight_uq_qr、weight_dkv_kr量化,默认值为0。
143 143 
144-- **kv_cache_quant_mode**(`int`):可选参数,表示kv_cache的量化模式。0表示非量化,1表示per-tensor量化,2表示per-channel量化,3-表示per-tile量化,默认值为0。144+- **kv_cache_quant_mode**(`int`):可选参数,表示kv_cache的量化模式。0表示非量化,1表示per-tensor量化,2表示per-channel量化,3-表示per-tile量化,默认值为0。
145 145 
146-- **query_quant_mode**(`int`):可选参数,表示query的量化模式。0表示非量化,1表示per-token-head量化,默认值为0。146+- **query_quant_mode**(`int`):可选参数,表示query的量化模式。0表示非量化,1表示per-token-head量化,默认值为0。
147 147 
148-- **ckvkr_repo_mode**(`int`):可选参数,表示kv_cache和kr_cache的存储模式。0表示kv_cache和kr_cache分别存储,1表示kv_cache和kr_cache合并存储,默认值为0。148+- **ckvkr_repo_mode**(`int`):可选参数,表示kv_cache和kr_cache的存储模式。0表示kv_cache和kr_cache分别存储,1表示kv_cache和kr_cache合并存储,默认值为0。
149 149 
150-- **quant_scale_repo_mode**(`int`):可选参数,表示量化scale的存储模式。0表示量化scale和数据分别存储,1表示量化scale和数据合并存储,默认值为0。150+- **quant_scale_repo_mode**(`int`):可选参数,表示量化scale的存储模式。0表示量化scale和数据分别存储,1表示量化scale和数据合并存储,默认值为0。
151 151 
152-- **tile_size**(`int`):可选参数,表示per-tile量化时每个tile的大小,仅在kv_cache_quant_mode为3时有效,默认值为128。152+- **tile_size**(`int`):可选参数,表示per-tile量化时每个tile的大小,仅在kv_cache_quant_mode为3时有效,默认值为128。
153 153 
154-- **qc_qr_scale**(`float`):可选参数,表示Query的尺度矫正系数,默认值为1.0。154+- **qc_qr_scale**(`float`):可选参数,表示Query的尺度矫正系数,默认值为1.0。
155 155 
156-- **kc_scale**(`float`):可选参数,表示Key的尺度矫正系数,默认值为1.0。156+- **kc_scale**(`float`):可选参数,表示Key的尺度矫正系数,默认值为1.0。
157 157 
158## 返回值说明158## 返回值说明
159-- **query**`Tensor`):表示Query的输出Tensor,即公式中q<sup>N</sup>。数据格式支持ND,dtype支持`bfloat16``int8`。shape支持3维和4维,格式为[T, N, Hckv]和[B, S, N, Hckv]。
160 159 
161-- **query_rope**(`Tensor`):表示Query位置编码的输出Tensor,即公式中q<sup>R</sup>。数据格式支持ND,dtype支持`bfloat16`。shape支持3维和4维,格式为[T, N, Dr]和[B, S, N, Dr]。160+- **query**(`Tensor`):表示Query的输出Tensor,即公式中q<sup>N</sup>。数据格式支持ND,dtype支持`bfloat16`和`int8`。shape支持3维和4维,格式为[T, N, Hckv]和[B, S, N, Hckv]。
162 161 
163-- **dequant_scale_q_nope**(`Tensor`):表示Query的输出Tensor的反量化参数。数据格式支持ND,dtype支持`float`。shape支持1维和3维,全量化kv_cache量化场景下,其shape为[T, N, 1]和[B*S, N, 1];其他场景下,其shape为[0]162+- **query_rope**(`Tensor`):表示Query位置编码的输出Tensor,即公式中q<sup>R</sup>。数据格式支持ND,dtype支持`bfloat16`。shape支持3维和4维,格式为[T, N, Dr]和[B, S, N, Dr]。
164 163 
165-- **query_norm**(`Tensor`):Query做RmsNorm_cq后的输出tensor(对应$q^C$)。数据格式支持ND,dtype支持`bfloat16`、`int8`。shape支持2维和3维,`query_norm_flag=True`时有效,shape为[T, Hcq][B, S, Hcq];`query_norm_flag=False`时无效,shape为[0]。164+- **dequant_scale_q_nope**(`Tensor`):表示Query的输出Tensor的反量化参数。数据格式支持ND,dtype支持`float`。shape支持1维和3维,全量化kv_cache量化场景下shape为[T, N, 1][B*S, N, 1];其他场景下shape为[0]。
166 165 
167-- **dequant_scale_q_norm**(`Tensor`):Query做RmsNorm_cq后的反量化参数。数据格式支持ND,数据类型支持`float`。shape支持1维和3维,`query_norm_flag=True`且`weight_quant_mode=1`或`weight_quant_mode=2`时有效,shape为[T, 1]或[B*S, 1];其余情况无效,shape为[0]。166+- **query_norm**(`Tensor`):Query做RmsNorm_cq后的输出tensor(对应$q^C$)。数据格式支持ND,dtype支持`bfloat16`、`int8`。shape支持2维和3维,`query_norm_flag=True`时有效,shape为[T, Hcq]或[B, S, Hcq];`query_norm_flag=False`时无效,shape为[0]。
167+ 
168+- **dequant_scale_q_norm**(`Tensor`):Query做RmsNorm_cq后的反量化参数。数据格式支持ND,数据类型支持`float`。shape支持1维和3维,`query_norm_flag=True`且`weight_quant_mode=1`或`weight_quant_mode=2`时有效,shape为[T, 1]或[B*S, 1];其余情况无效,shape为[0]。
168 169 
169## 约束说明170## 约束说明
171+ 
170- 该接口支持推理场景下使用。172- 该接口支持推理场景下使用。
171 173 
172- 该接口支持单算子模式和图模式。174- 该接口支持单算子模式和图模式。
173 175 
174-- shape 格式字段含义说明176+- shape 格式字段含义说明
177+ 
175 | 字段名 | 英文全称/含义 | 取值规则与说明 |178 | 字段名 | 英文全称/含义 | 取值规则与说明 |
176 |--------------|--------------------------------|------------------------------------------------------------------------------|179 |--------------|--------------------------------|------------------------------------------------------------------------------|
177 | B | Batch(输入样本批量大小) | 取值范围:0~65536 |180 | B | Batch(输入样本批量大小) | 取值范围:0~65536 |
@@ -540,10 +543,9 @@ torch_npu.npu_mla_prolog_v3(token_x, weight_dq, weight_uq_qr, weight_uk, weight_
540 </table>543 </table>
541 </div>544 </div>
542 545 
543- 
544## 调用示例546## 调用示例
545 547 
546-- 单算子模式调用548+- 单算子模式调用
547 549 
548 ```python550 ```python
549 import torch551 import torch
@@ -596,8 +598,7 @@ torch_npu.npu_mla_prolog_v3(token_x, weight_dq, weight_uq_qr, weight_uk, weight_
596 device='npu:0', dtype=torch.bfloat16)598 device='npu:0', dtype=torch.bfloat16)
597 ```599 ```
598 600 
599- 601+- 图模式调用
600-- 图模式调用
601 602 
602 ```python603 ```python
603 # 入图方式604 # 入图方式
@@ -687,4 +688,4 @@ torch_npu.npu_mla_prolog_v3(token_x, weight_dq, weight_uq_qr, weight_uk, weight_
687 [ 0.0180, 0.0186, -0.0067, ..., 0.0204, -0.0045, -0.0164],688 [ 0.0180, 0.0186, -0.0067, ..., 0.0204, -0.0045, -0.0164],
688 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],689 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],
689 device='npu:0', dtype=torch.bfloat16) 690 device='npu:0', dtype=torch.bfloat16)
690- ```691+ ```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_mm_all_reduce_base.md+6-2
@@ -20,7 +20,7 @@
20 20 
21## 函数原型21## 函数原型
22 22 
23-```23+```python
24torch_npu.npu_mm_all_reduce_base(x1, x2, hcom, *, reduce_op='sum', bias=None, antiquant_scale=None, antiquant_offset=None, x3=None, dequant_scale=None, pertoken_scale=None, comm_quant_scale_1=None, comm_quant_scale_2=None, comm_turn=0, antiquant_group_size=0) -> Tensor24torch_npu.npu_mm_all_reduce_base(x1, x2, hcom, *, reduce_op='sum', bias=None, antiquant_scale=None, antiquant_offset=None, x3=None, dequant_scale=None, pertoken_scale=None, comm_quant_scale_1=None, comm_quant_scale_2=None, comm_turn=0, antiquant_group_size=0) -> Tensor
25```25```
26 26 
@@ -55,6 +55,7 @@ torch_npu.npu_mm_all_reduce_base(x1, x2, hcom, *, reduce_op='sum', bias=None, an
55- **antiquant_group_size** (`int`):可选参数。表示伪量化pergroup算法模式下,对输入`x2`进行反量化计算的groupSize输入,描述一组反量化参数对应的待反量化数据量在$k$轴方向的大小。当伪量化算法模式不为pergroup时传入`0`;当伪量化算法模式为pergroup时传入值的范围为`[32, min(k-1, INT_MAX)]`且值要求是32的倍数,其中$k$为`x2`第一维的大小。默认值`0`,为`0`则表示非pergroup场景。55- **antiquant_group_size** (`int`):可选参数。表示伪量化pergroup算法模式下,对输入`x2`进行反量化计算的groupSize输入,描述一组反量化参数对应的待反量化数据量在$k$轴方向的大小。当伪量化算法模式不为pergroup时传入`0`;当伪量化算法模式为pergroup时传入值的范围为`[32, min(k-1, INT_MAX)]`且值要求是32的倍数,其中$k$为`x2`第一维的大小。默认值`0`,为`0`则表示非pergroup场景。
56 56 
57## 返回值说明57## 返回值说明
58+ 
58`Tensor`59`Tensor`
59 60 
60数据类型非量化场景以及伪量化场景与`x1`保持一致,全量化场景输出数据类型为`float16``bfloat16`。shape第0维度和`x1`的第0维保持一致,若`x1`为2维,shape第1维度和`x2`的第1维保持一致,若`x1`为3维,shape第1维度和`x1`的第1维保持一致,shape第2维度和`x2`的第1维保持一致。61数据类型非量化场景以及伪量化场景与`x1`保持一致,全量化场景输出数据类型为`float16``bfloat16`。shape第0维度和`x1`的第0维保持一致,若`x1`为2维,shape第1维度和`x2`的第1维保持一致,若`x1`为3维,shape第1维度和`x1`的第1维保持一致,shape第2维度和`x2`的第1维保持一致。
@@ -81,18 +82,21 @@ torch_npu.npu_mm_all_reduce_base(x1, x2, hcom, *, reduce_op='sum', bias=None, an
81- 不同场景下数据类型支持情况:82- 不同场景下数据类型支持情况:
82 83 
83 **表1** 非量化场景84 **表1** 非量化场景
85+ 
84 |产品型号|x1|x2|bias|x3|output(输出)|antiquant_scale|antiquant_offset|dequant_scale|86 |产品型号|x1|x2|bias|x3|output(输出)|antiquant_scale|antiquant_offset|dequant_scale|
85 |--------|--------|--------|--------|--------|--------|--------|--------|--------|87 |--------|--------|--------|--------|--------|--------|--------|--------|--------|
86 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`float16`|`float16`|`float16`|`float16`|`float16`|None|None|None|88 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`float16`|`float16`|`float16`|`float16`|`float16`|None|None|None|
87 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|None|None|None|89 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|None|None|None|
88 90
89 **表2** 伪量化场景91 **表2** 伪量化场景
92+ 
90 |产品型号|x1|x2|bias|x3|output(输出)|antiquant_scale|antiquant_offset|dequant_scale|93 |产品型号|x1|x2|bias|x3|output(输出)|antiquant_scale|antiquant_offset|dequant_scale|
91 |--------|--------|--------|--------|--------|--------|--------|--------|--------|94 |--------|--------|--------|--------|--------|--------|--------|--------|--------|
92 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`float16`|`int8`|`float16`|`float16`|`float16`|`float16`|`float16`|None|95 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`float16`|`int8`|`float16`|`float16`|`float16`|`float16`|`float16`|None|
93 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`bfloat16`|`int8`|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|None|96 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`bfloat16`|`int8`|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|`bfloat16`|None|
94 97
95 **表3** 全量化场景98 **表3** 全量化场景
99+ 
96 |产品型号|x1|x2|bias|x3|output(输出)|antiquant_scale|antiquant_offset|dequant_scale|pertoken_scale|100 |产品型号|x1|x2|bias|x3|output(输出)|antiquant_scale|antiquant_offset|dequant_scale|pertoken_scale|
97 |--------|--------|--------|--------|--------|--------|--------|--------|--------|--------|101 |--------|--------|--------|--------|--------|--------|--------|--------|--------|--------|
98 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`int8`|`int8`|`int32`|`float16`|`float16`|None|None|`uint64`或`int64`|None|102 |Atlas A2 训练系列产品/Atlas A2 推理系列产品|`int8`|`int8`|`int32`|`float16`|`float16`|None|None|`uint64`或`int64`|None|
@@ -379,4 +383,4 @@ torch_npu.npu_mm_all_reduce_base(x1, x2, hcom, *, reduce_op='sum', bias=None, an
379 6.2500e-02, -2.0550e+02],383 6.2500e-02, -2.0550e+02],
380 [ 7.4062e+01, -6.0100e+02, -3.0750e+02, ..., -2.1500e+02,384 [ 7.4062e+01, -6.0100e+02, -3.0750e+02, ..., -2.1500e+02,
381 -2.4450e+02, 3.2400e+02]], device='npu:0', dtype=torch.float16)385 -2.4450e+02, 3.2400e+02]], device='npu:0', dtype=torch.float16)
382- ```386+ ```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_mm_reduce_scatter_base.md+15-15
@@ -9,10 +9,10 @@
9 9 
10## 功能说明10## 功能说明
11 11 
12-- API功能:TP切分场景下,实现matmul和reduce\_scatter的融合,融合算子内部实现计算和通信流水并行。支持perchannel,pertoken量化。12+- API功能:TP切分场景下,实现matmul和reduce\_scatter的融合,融合算子内部实现计算和通信流水并行。支持perchannel,pertoken量化。
13 13 
14-- 计算公式:14+- 计算公式:
15- $x1$代表输入`input`15+ $x1$代表输入`input`
16 16
17 基础场景:17 基础场景:
18 $$18 $$
@@ -28,7 +28,7 @@
28 28 
29## 函数原型29## 函数原型
30 30 
31-```31+```python
32torch_npu.npu_mm_reduce_scatter_base(input, x2, hcom, world_size, *, reduce_op='sum', bias=None, x1_scale=None, x2_scale=None, comm_turn=0, output_dtype=None, comm_mode=None) -> Tensor32torch_npu.npu_mm_reduce_scatter_base(input, x2, hcom, world_size, *, reduce_op='sum', bias=None, x1_scale=None, x2_scale=None, comm_turn=0, output_dtype=None, comm_mode=None) -> Tensor
33```33```
34 34 
@@ -38,8 +38,8 @@ torch_npu.npu_mm_reduce_scatter_base(input, x2, hcom, world_size, *, reduce_op='
38- **x2** (`Tensor`):必选参数。数据类型与`input`一致,数据格式支持$ND$、$NZ$。$NZ$仅在`comm_mode``aiv`时支持。输入shape支持2维,形如\(k, n\)。轴满足matmul算子入参要求,k轴相等,且k轴取值范围为\[256, 65535\),m轴需要整除`world_size`38- **x2** (`Tensor`):必选参数。数据类型与`input`一致,数据格式支持$ND$、$NZ$。$NZ$仅在`comm_mode``aiv`时支持。输入shape支持2维,形如\(k, n\)。轴满足matmul算子入参要求,k轴相等,且k轴取值范围为\[256, 65535\),m轴需要整除`world_size`
39- **hcom** (`str`):必选参数。通信域handle名,通过get\_hccl\_comm\_name接口获取。39- **hcom** (`str`):必选参数。通信域handle名,通过get\_hccl\_comm\_name接口获取。
40- **world\_size** (`int`):必选参数。通信域内的rank总数。40- **world\_size** (`int`):必选参数。通信域内的rank总数。
41- - <term>Atlas A2 训练系列产品</term>支持2、4、8卡,支持HCCS链路all mesh组网(每张卡和其它卡两两相连)。41+ - <term>Atlas A2 训练系列产品</term>支持2、4、8卡,支持HCCS链路all mesh组网(每张卡和其它卡两两相连)。
42- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>支持2、4、8、16、32卡,支持HCCS链路double ring组网(多张卡按顺序组成一个圈,每张卡只和左右卡相连)。42+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>支持2、4、8、16、32卡,支持HCCS链路double ring组网(多张卡按顺序组成一个圈,每张卡只和左右卡相连)。
43 43 
44- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。44- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
45- **reduce\_op** (`str`):可选参数。reduce操作类型,当前仅支持'sum',默认值为'sum'。45- **reduce\_op** (`str`):可选参数。reduce操作类型,当前仅支持'sum',默认值为'sum'。
@@ -59,16 +59,17 @@ shape维度和`input`保持一致。
59量化场景下,`x2_scale``int64`数据类型时,输出数据类型为`float16``x1_scale``x2_scale`均为`float32`时, 输出数据类型由`output_dtype`指定,默认为`bfloat16`59量化场景下,`x2_scale``int64`数据类型时,输出数据类型为`float16``x1_scale``x2_scale`均为`float32`时, 输出数据类型由`output_dtype`指定,默认为`bfloat16`
60 60 
61## 约束说明61## 约束说明
62-- `input`不支持输入转置后的tensor,`x2`转置后输入,需要满足shape的第一维大小与`input`的最后一维相同,满足matmul的计算条件。62+ 
63-- `comm_mode``ai_cpu`时:63+- `input`不支持输入转置后的tensor,`x2`转置后输入,需要满足shape的第一维大小与`input`的最后一维相同,满足matmul的计算条件。
64- - 该接口仅在训练场景下使用。64+- `comm_mode`为`ai_cpu`时:
65- - 该接口支持图模式65+ - 该接口仅在训练场景下使用
66- - <term>Atlas A2 训练系列产品</term>:一个模型中的通算融合算子(AllGatherMatmul、MatmulReduceScatter、MatmulAllReduce),仅支持相同通信域66+ - 该接口支持图模式
67-- `comm_mode`为`aiv`时,训练和推理场景均可使用 67+ - <term>Atlas A2 训练系列产品</term>:一个模型中的通算融合算子(AllGatherMatmul、MatmulReduceScatter、MatmulAllReduce),仅支持相同通信域
68+- `comm_mode``aiv`时,训练和推理场景均可使用。
68 69 
69## 调用示例70## 调用示例
70 71 
71-- 单算子模式调用72+- 单算子模式调用
72 73 
73 ```python74 ```python
74 import torch75 import torch
@@ -101,7 +102,7 @@ shape维度和`input`保持一致。
101 mp.spawn(run_mm_reduce_scatter_base, args=(worksize, master_ip, master_port, x1_shape, x2_shape, dtype), nprocs=worksize)102 mp.spawn(run_mm_reduce_scatter_base, args=(worksize, master_ip, master_port, x1_shape, x2_shape, dtype), nprocs=worksize)
102 ```103 ```
103 104 
104-- 图模式调用105+- 图模式调用
105 106 
106 ```python107 ```python
107 import torch108 import torch
@@ -154,4 +155,3 @@ shape维度和`input`保持一致。
154 dtype = torch.float16155 dtype = torch.float16
155 mp.spawn(run_mm_reduce_scatter_base, args=(worksize, master_ip, master_port, x1_shape, x2_shape, dtype), nprocs=worksize)156 mp.spawn(run_mm_reduce_scatter_base, args=(worksize, master_ip, master_port, x1_shape, x2_shape, dtype), nprocs=worksize)
156 ```157 ```
157- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_compute_expert_tokens.md+2-2
@@ -17,7 +17,7 @@
17 17 
18## 函数原型18## 函数原型
19 19 
20-```20+```python
21torch_npu.npu_moe_compute_expert_tokens(sorted_expert_for_source_row, num_expert) -> Tensor21torch_npu.npu_moe_compute_expert_tokens(sorted_expert_for_source_row, num_expert) -> Tensor
22```22```
23 23 
@@ -28,6 +28,7 @@ torch_npu.npu_moe_compute_expert_tokens(sorted_expert_for_source_row, num_expert
28- **num_expert** (`int`):必选参数,表示总专家数。对应公式中的$numExpert$。28- **num_expert** (`int`):必选参数,表示总专家数。对应公式中的$numExpert$。
29 29 
30## 返回值说明30## 返回值说明
31+ 
31`Tensor`32`Tensor`
32 33 
33 对应公式中的$expertTokens$,要求的是一个1维张量,数据类型与`sorted_expert_for_source_row`保持一致。34 对应公式中的$expertTokens$,要求的是一个1维张量,数据类型与`sorted_expert_for_source_row`保持一致。
@@ -74,4 +75,3 @@ torch_npu.npu_moe_compute_expert_tokens(sorted_expert_for_source_row, num_expert
74 if __name__ == '__main__':75 if __name__ == '__main__':
75 main()76 main()
76 ```77 ```
77- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_distribute_combine.md+95-95
@@ -9,9 +9,9 @@
9 9 
10## 功能说明<a name="zh-cn_topic_0000002168254826_section14441124184110"></a>10## 功能说明<a name="zh-cn_topic_0000002168254826_section14441124184110"></a>
11 11 
12-- API功能:先进行reduce\_scatterv通信,再进行alltoallv通信,最后将接收的数据整合(乘权重再相加)。需与[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)配套使用,相当于按npu\_moe\_distribute\_dispatch算子收集数据的路径原路返回。12+- API功能:先进行reduce\_scatterv通信,再进行alltoallv通信,最后将接收的数据整合(乘权重再相加)。需与[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)配套使用,相当于按npu\_moe\_distribute\_dispatch算子收集数据的路径原路返回。
13 13 
14-- 计算公式:14+- 计算公式:
15 15 
16 $$16 $$
17 rs\_out = ReduceScatterV(expend\_x)\\17 rs\_out = ReduceScatterV(expend\_x)\\
@@ -22,147 +22,148 @@
22 22 
23## 函数原型<a name="zh-cn_topic_0000002168254826_section45077510411"></a>23## 函数原型<a name="zh-cn_topic_0000002168254826_section45077510411"></a>
24 24 
25-```25+```python
26torch_npu.npu_moe_distribute_combine(expand_x, expert_ids, expand_idx, ep_send_counts, expert_scales, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, tp_send_counts=None, x_active_mask=None, activation_scale=None, weight_scale=None, group_list=None, expand_scales=None, shared_expert_x=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, global_bs=0, out_dtype=0, comm_quant_mode=0, group_list_type=0) -> Tensor26torch_npu.npu_moe_distribute_combine(expand_x, expert_ids, expand_idx, ep_send_counts, expert_scales, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, tp_send_counts=None, x_active_mask=None, activation_scale=None, weight_scale=None, group_list=None, expand_scales=None, shared_expert_x=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, global_bs=0, out_dtype=0, comm_quant_mode=0, group_list_type=0) -> Tensor
27```27```
28 28 
29## 参数说明<a name="zh-cn_topic_0000002168254826_section112637109429"></a>29## 参数说明<a name="zh-cn_topic_0000002168254826_section112637109429"></a>
30 30 
31-- **expand\_x** (`Tensor`):必选参数。根据`expert_ids`进行扩展过的token特征,要求为2维张量,shape为\(max\(tp\_world\_size, 1\) \*A, H\),数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。31+- **expand\_x** (`Tensor`):必选参数。根据`expert_ids`进行扩展过的token特征,要求为2维张量,shape为\(max\(tp\_world\_size, 1\) \*A, H\),数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。
32- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家场景。32+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家场景。
33 33 
34-- **expert\_ids** (`Tensor`):必选参数。每个token的topK个专家索引,要求为2维张量,shape为\(BS, K\)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`expert_ids`输入,张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。34+- **expert\_ids** (`Tensor`):必选参数。每个token的topK个专家索引,要求为2维张量,shape为\(BS, K\)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`expert_ids`输入,张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。
35-- **expand\_idx** (`Tensor`):必选参数。表示给同一专家发送的token个数,要求为1维张量。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`expand_idx`输出。35+- **expand\_idx** (`Tensor`):必选参数。表示给同一专家发送的token个数,要求为1维张量。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`expand_idx`输出。
36- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(BS \* K,\)。36+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(BS \* K,\)。
37- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(BS \* K,\)。37+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(BS \* K,\)。
38 38 
39-- **ep\_send\_counts** (`Tensor`):必选参数。示本卡每个专家发给EP(Expert Parallelism)域每个卡的token数(token数以前缀和的形式表示),要求为1维张量。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`ep_recv_counts`输出。39+- **ep\_send\_counts** (`Tensor`):必选参数。示本卡每个专家发给EP(Expert Parallelism)域每个卡的token数(token数以前缀和的形式表示),要求为1维张量。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`ep_recv_counts`输出。
40- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(moe\_expert\_num+2\*global\_bs\*K\*server\_num,\),前`moe_expert_num`个数表示在EP通信域内,该卡上每个专家收到来自其他各卡的token数(以前缀和的形式表示),2\*global\_bs\*K\*server\_num用于存储机间和机内通信前,combine可提前做reduce操作的token个数和通信区偏移量,`global_bs`传入0时此处按照bs\*ep\_world\_size计算。40+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(moe\_expert\_num+2\*global\_bs\*K\*server\_num,\),前`moe_expert_num`个数表示在EP通信域内,该卡上每个专家收到来自其他各卡的token数(以前缀和的形式表示),2\*global\_bs\*K\*server\_num用于存储机间和机内通信前,combine可提前做reduce操作的token个数和通信区偏移量,`global_bs`传入0时此处按照bs\*ep\_world\_size计算。
41- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(ep\_world\_size\*max\(tp\_world\_size, 1\)\*local\_expert\_num,\)。41+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(ep\_world\_size\*max\(tp\_world\_size, 1\)\*local\_expert\_num,\)。
42 42 
43-- **expert\_scales** (`Tensor`):必选参数。表示每个token的topK个专家的权重,要求为2维张量,shape为\(BS, K\),其中共享专家不需要乘权重系数,直接相加即可。数据类型支持`float`,数据格式为$ND$,支持非连续的Tensor。43+- **expert\_scales** (`Tensor`):必选参数。表示每个token的topK个专家的权重,要求为2维张量,shape为\(BS, K\),其中共享专家不需要乘权重系数,直接相加即可。数据类型支持`float`,数据格式为$ND$,支持非连续的Tensor。
44-- **group\_ep** (`str`):必选参数。EP通信域名称,专家并行的通信域。字符串长度范围为\[1, 128\)。<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>时不能和`group_tp`相同。44+- **group\_ep** (`str`):必选参数。EP通信域名称,专家并行的通信域。字符串长度范围为\[1, 128\)。<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>时不能和`group_tp`相同。
45-- **ep\_world\_size** (`int`):必选参数,EP通信域size。45+- **ep\_world\_size** (`int`):必选参数,EP通信域size。
46- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值支持16、32、64。46+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值支持16、32、64。
47- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持8、16、32、64、128、144、256、288。47+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持8、16、32、64、128、144、256、288。
48 48 
49-- **ep\_rank\_id** (`int`):必选参数,EP通信域本卡ID,取值范围\[0, ep\_world\_size\),同一个EP通信域中各卡的`ep_rank_id`不重复。49+- **ep\_rank\_id** (`int`):必选参数,EP通信域本卡ID,取值范围\[0, ep\_world\_size\),同一个EP通信域中各卡的`ep_rank_id`不重复。
50-- **moe\_expert\_num** (`int`):必选参数,MoE专家数量,取值范围\[1, 512\],并且满足以下条件:moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0。50+- **moe\_expert\_num** (`int`):必选参数,MoE专家数量,取值范围\[1, 512\],并且满足以下条件:moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0。
51- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:还需满足moe\_expert\_num\/\(ep\_world\_size - shared\_expert\_rank\_num\) <= 24。51+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:还需满足moe\_expert\_num\/\(ep\_world\_size - shared\_expert\_rank\_num\) <= 24。
52- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。52- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
53-- **tp\_send\_counts** (`Tensor`):可选参数,表示本卡每个专家发给TP(Tensor Parallelism)通信域每个卡的数据量。对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`tp_recv_counts`输出。53+- **tp\_send\_counts** (`Tensor`):可选参数,表示本卡每个专家发给TP(Tensor Parallelism)通信域每个卡的数据量。对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`tp_recv_counts`输出。
54- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP通信域,使用默认输入None。54+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP通信域,使用默认输入None。
55- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求为一个1维张量,shape为\(tp\_world\_size,\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。55+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求为一个1维张量,shape为\(tp\_world\_size,\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。
56 56 
57-- **x\_active\_mask** (`Tensor`):预留参数,暂未使用,使用默认值即可。57+- **x\_active\_mask** (`Tensor`):预留参数,暂未使用,使用默认值即可。
58-- **activation\_scale** (`Tensor`):预留参数,暂未使用,使用默认值即可。58+- **activation\_scale** (`Tensor`):预留参数,暂未使用,使用默认值即可。
59-- **weight\_scale** (`Tensor`):预留参数,暂未使用,使用默认值即可。59+- **weight\_scale** (`Tensor`):预留参数,暂未使用,使用默认值即可。
60-- **group\_list** (`Tensor`):预留参数,暂未使用,使用默认值即可。60+- **group\_list** (`Tensor`):预留参数,暂未使用,使用默认值即可。
61-- **expand\_scales** (`Tensor`):对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`expand_scales`输出。61+- **expand\_scales** (`Tensor`):对应[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)的`expand_scales`输出。
62- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:必选参数,要求为1维张量,shape为\(A,\),数据类型支持`float`,数据格式为$ND$,支持非连续的Tensor。62+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:必选参数,要求为1维张量,shape为\(A,\),数据类型支持`float`,数据格式为$ND$,支持非连续的Tensor。
63- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:暂不支持该参数,使用默认值即可。63+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:暂不支持该参数,使用默认值即可。
64 64 
65-- **shared\_expert\_x** (`Tensor`):预留参数,暂未使用,使用默认值即可。65+- **shared\_expert\_x** (`Tensor`):预留参数,暂未使用,使用默认值即可。
66 66 
67-- **group\_tp** (`str`):可选参数,TP通信域名称,数据并行的通信域。有TP域通信才需要传参,若无TP域通信,使用默认值""即可。67+- **group\_tp** (`str`):可选参数,TP通信域名称,数据并行的通信域。有TP域通信才需要传参,若无TP域通信,使用默认值""即可。
68- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:eager模式使用默认值即可,图模式传入与`group_ep`相同。68+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:eager模式使用默认值即可,图模式传入与`group_ep`相同。
69- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。69+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。
70 70 
71-- **tp\_world\_size** (`int`):可选参数,TP通信域size。有TP域通信才需要传参。71+- **tp\_world\_size** (`int`):可选参数,TP通信域size。有TP域通信才需要传参。
72- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。72+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。
73- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。73+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。
74 74 
75-- **tp\_rank\_id** (`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。75+- **tp\_rank\_id** (`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。
76- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。76+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。
77- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],同一个TP通信域中各卡的`tp_rank_id`不重复。无TP域通信时,传0即可。77+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],同一个TP通信域中各卡的`tp_rank_id`不重复。无TP域通信时,传0即可。
78 78 
79-- **expert\_shard\_type** (`int`):表示共享专家卡排布类型。79+- **expert\_shard\_type** (`int`):表示共享专家卡排布类型。
80- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。80+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。
81- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。81+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。
82 82 
83-- **shared\_expert\_num** (`int`):表示共享专家数量,一个共享专家可以复制部署到多个卡上。83+- **shared\_expert\_num** (`int`):表示共享专家数量,一个共享专家可以复制部署到多个卡上。
84- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。84+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。
85- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅支持1,默认值为1。85+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅支持1,默认值为1。
86 86 
87-- **shared\_expert\_rank\_num** (`int`):可选参数,表示共享专家卡数量。87+- **shared\_expert\_rank\_num** (`int`):可选参数,表示共享专家卡数量。
88- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家,传0即可。88+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家,传0即可。
89- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, ep\_world\_size\)。取0表示无共享专家,不取0时需满足ep\_world\_size%shared\_expert\_rank\_num=0。89+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, ep\_world\_size\)。取0表示无共享专家,不取0时需满足ep\_world\_size%shared\_expert\_rank\_num=0。
90 90 
91-- **global\_bs** (`int`):可选参数,表示EP域全局的BS(batch size)大小。91+- **global\_bs** (`int`):可选参数,表示EP域全局的BS(batch size)大小。
92- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size或者256\*ep\_world\_size,其中max\_bs表示单rank BS最大值,建议按max\_bs\*ep\_world\_size传入,固定按256\*ep\_world\_size传入,在后续版本BS大于256的场景下会无法支持;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。92+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size或者256\*ep\_world\_size,其中max\_bs表示单rank BS最大值,建议按max\_bs\*ep\_world\_size传入,固定按256\*ep\_world\_size传入,在后续版本BS大于256的场景下会无法支持;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。
93- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。93+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。
94 94 
95-- **out\_dtype** (`int`):预留参数,暂未使用,使用默认值即可。95+- **out\_dtype** (`int`):预留参数,暂未使用,使用默认值即可。
96-- **comm\_quant\_mode** (`int`):表示通信量化类型。96+- **comm\_quant\_mode** (`int`):表示通信量化类型。
97- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:支持取0和2。0表示通信时不量化,2表示通信时进行`int8`量化。仅当HCCL\_INTRA\_PCIE\_ENABLE=1且HCCL\_INTRA\_ROCE\_ENABLE=0且驱动版本不低于25.0.RC1.1时才支持取2。97+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:支持取0和2。0表示通信时不量化,2表示通信时进行`int8`量化。仅当HCCL\_INTRA\_PCIE\_ENABLE=1且HCCL\_INTRA\_ROCE\_ENABLE=0且驱动版本不低于25.0.RC1.1时才支持取2。
98- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持取0和2。0表示通信时不量化,2表示通信时进行`int8`量化。当且仅当`tp_world_size`不等于2时,可以使能`int8`量化。98+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持取0和2。0表示通信时不量化,2表示通信时进行`int8`量化。当且仅当`tp_world_size`不等于2时,可以使能`int8`量化。
99 99 
100-- **group\_list\_type** (`int`):预留参数,暂未使用,使用默认值即可。100+- **group\_list\_type** (`int`):预留参数,暂未使用,使用默认值即可。
101 101 
102## 返回值说明<a name="zh-cn_topic_0000002168254826_section22231435517"></a>102## 返回值说明<a name="zh-cn_topic_0000002168254826_section22231435517"></a>
103+ 
103`Tensor`104`Tensor`
104 105 
105表示处理后的token,要求为2维张量,shape为\(BS, H\),数据类型支持`bfloat16``float16`,类型与输入`expand_x`保持一致,数据格式为$ND$,不支持非连续的Tensor。106表示处理后的token,要求为2维张量,shape为\(BS, H\),数据类型支持`bfloat16``float16`,类型与输入`expand_x`保持一致,数据格式为$ND$,不支持非连续的Tensor。
106 107 
107## 约束说明<a name="zh-cn_topic_0000002168254826_section12345537164214"></a>108## 约束说明<a name="zh-cn_topic_0000002168254826_section12345537164214"></a>
108 109 
109-- 该接口支持推理场景下使用。110+- 该接口支持推理场景下使用。
110-- 该接口支持静态图模式,`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`必须配套使用。111+- 该接口支持静态图模式,`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`必须配套使用。
111-- 在不同产品型号、不同通信算法或不同版本中,`npu_moe_distribute_dispatch`的Tensor输出`expand_idx`、`ep_recv_counts`、`tp_recv_counts`、`expand_scales`中的元素值可能不同,使用时直接将上述Tensor传给`npu_moe_distribute_combine`对应参数即可,模型其他业务逻辑不应对其存在依赖。112+- 在不同产品型号、不同通信算法或不同版本中,`npu_moe_distribute_dispatch`的Tensor输出`expand_idx`、`ep_recv_counts`、`tp_recv_counts`、`expand_scales`中的元素值可能不同,使用时直接将上述Tensor传给`npu_moe_distribute_combine`对应参数即可,模型其他业务逻辑不应对其存在依赖。
112-- 调用接口过程中使用的`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`shared_expert_num`、`shared_expert_rank_num`、`global_bs`参数取值所有卡需保持一致,`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp、tp_world_size`、`expert_shard_type`、`global_bs`网络中不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)对应参数也保持一致。113+- 调用接口过程中使用的`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`shared_expert_num`、`shared_expert_rank_num`、`global_bs`参数取值所有卡需保持一致,`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp、tp_world_size`、`expert_shard_type`、`global_bs`网络中不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_dispatch](torch_npu-npu_moe_distribute_dispatch.md)对应参数也保持一致。
113-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。114+- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。
114-- 参数里Shape使用的变量如下:115+- 参数里Shape使用的变量如下:
115- - A:表示本卡接收的最大token数量,取值范围如下:116+ - A:表示本卡接收的最大token数量,取值范围如下:
116- - 对于共享专家,当`global_bs`为0时,要满足A=BS\*shared\_expert\_num/shared\_expert\_rank\_num;当`global_bs`非0时,要满足A=global\_bs\*shared\_expert\_num/shared\_expert\_rank\_num。117+ - 对于共享专家,当`global_bs`为0时,要满足A=BS\*shared\_expert\_num/shared\_expert\_rank\_num;当`global_bs`非0时,要满足A=global\_bs\*shared\_expert\_num/shared\_expert\_rank\_num。
117- - 对于MoE专家,当`global_bs`为0时,要满足A\>=BS\*ep\_world\_size\*min\(local\_expert\_num, K\);当`global_bs`非0时,要满足A\>=global\_bs\* min\(local\_expert\_num, K\)。118+ - 对于MoE专家,当`global_bs`为0时,要满足A\>=BS\*ep\_world\_size\*min\(local\_expert\_num, K\);当`global_bs`非0时,要满足A\>=global\_bs\* min\(local\_expert\_num, K\)。
118 119 
119- - H:表示hidden size隐藏层大小。120+ - H:表示hidden size隐藏层大小。
120- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值范围\(0, 7168\],且保证是32的整数倍。121+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值范围\(0, 7168\],且保证是32的整数倍。
121- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅支持 7168。122+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅支持 7168。
122 123 
123- - BS:表示待发送的token数量。124+ - BS:表示待发送的token数量。
124- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值范围为0<BS≤256。125+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值范围为0<BS≤256。
125- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为0<BS≤512。126+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为0<BS≤512。
126 127 
127- - K:表示选取topK个专家,需满足0<K≤moe\_expert\_num。128+ - K:表示选取topK个专家,需满足0<K≤moe\_expert\_num。
128- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:保证取值范围为0<K≤16。129+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:保证取值范围为0<K≤16。
129- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:保证取值范围为0<K≤8。130+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:保证取值范围为0<K≤8。
130 131 
131- - server\_num:表示服务器的节点数,取值只支持2、4、8。132+ - server\_num:表示服务器的节点数,取值只支持2、4、8。
132- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:仅该场景的shape使用了该变量。133+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:仅该场景的shape使用了该变量。
133 134 
134- - local\_expert\_num:表示本卡专家数量。135+ - local\_expert\_num:表示本卡专家数量。
135- - 对于共享专家卡,local\_expert\_num=1136+ - 对于共享专家卡,local\_expert\_num=1
136- - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num),当local\_expert\_num\>1时,不支持TP域通信。137+ - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num),当local\_expert\_num\>1时,不支持TP域通信。
137 138 
138-- HCCL通信域缓存区大小:139+- HCCL通信域缓存区大小:
139 140 
140 调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。141 调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。
141- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:142+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
142 该场景支持通过环境变量HCCL\_BUFFSIZE配置。143 该场景支持通过环境变量HCCL\_BUFFSIZE配置。
143 - 设置大小要求\>=2\*\(BS\*ep\_world\_size\*min\(local\_expert\_num, K\)\*H\*sizeof\(uint16\)+2MB\)。144 - 设置大小要求\>=2\*\(BS\*ep\_world\_size\*min\(local\_expert\_num, K\)\*H\*sizeof\(uint16\)+2MB\)。
144- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:145+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
145 该场景不仅支持通过环境变量HCCL\_BUFFSIZE配置,还支持通过hccl_buffer_size配置(参考《[PyTorch训练模型迁移调优](https://hiascend.com/document/redirect/canncommercial-ptmigr)》中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。146 该场景不仅支持通过环境变量HCCL\_BUFFSIZE配置,还支持通过hccl_buffer_size配置(参考《[PyTorch训练模型迁移调优](https://hiascend.com/document/redirect/canncommercial-ptmigr)》中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。
146 - ep通信域内:设置大小要求\>=2且满足1024\^2\*\(HCCL\_BUFFSIZE\-2\)\/2\>=BS\*2\*\(H\+128\)\*\(ep\_world\_size\*local\_expert\_num\+K\+1\),local\_expert\_num需使用MoE专家卡的本卡专家数。147 - ep通信域内:设置大小要求\>=2且满足1024\^2\*\(HCCL\_BUFFSIZE\-2\)\/2\>=BS\*2\*\(H\+128\)\*\(ep\_world\_size\*local\_expert\_num\+K\+1\),local\_expert\_num需使用MoE专家卡的本卡专家数。
147- - tp通信域内:设置大小要求\>=A * (H * 2 + 128) * 2。148+ - tp通信域内:设置大小要求\>=A \* (H \* 2 + 128) * 2。
148-- <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:配置环境变量HCCL\_INTRA\_PCIE\_ENABLE=1和HCCL\_INTRA\_ROCE\_ENABLE=0可以减少跨机通信数据量,提升算子性能。此时要求HCCL\_BUFFSIZE\>=moe\_expert\_num\*BS\*\(H\*sizeof\(dtype_x\)+4\*\(\(K+7\)/8\*8\)\*sizeof\(uint32\)\)+4MB+100MB。并且,对于入参moe\_expert\_num,只要求moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0,不要求moe\_expert\_num\/\(ep\_world\_size - shared\_expert\_rank\_num\) <= 24。149+- <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:配置环境变量HCCL\_INTRA\_PCIE\_ENABLE=1和HCCL\_INTRA\_ROCE\_ENABLE=0可以减少跨机通信数据量,提升算子性能。此时要求HCCL\_BUFFSIZE\>=moe\_expert\_num\*BS\*\(H\*sizeof\(dtype_x\)+4\*\(\(K+7\)/8\*8\)\*sizeof\(uint32\)\)+4MB+100MB。并且,对于入参moe\_expert\_num,只要求moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0,不要求moe\_expert\_num\/\(ep\_world\_size - shared\_expert\_rank\_num\) <= 24。
149 150 
150-- 本文公式中的“/”表示整除。151+- 本文公式中的“/”表示整除。
151 152 
152-- 通信域使用约束:153+- 通信域使用约束:
153 154 
154- - 一个模型中的`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`算子仅支持相同EP通信域,且该通信域中不允许有其他算子。155+ - 一个模型中的`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`算子仅支持相同EP通信域,且该通信域中不允许有其他算子。
155 156 
156- - 一个模型中的`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。157+ - 一个模型中的`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。
157 158 
158- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:一个通信域内的节点需在一个超节点内,不支持跨超节点。159+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:一个通信域内的节点需在一个超节点内,不支持跨超节点。
159 160 
160-- 组网约束:161+- 组网约束:
161- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:多机场景仅支持交换机组网,不支持双机直连组网。162+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:多机场景仅支持交换机组网,不支持双机直连组网。
162 163 
163## 调用示例<a name="zh-cn_topic_0000002168254826_section14459801435"></a>164## 调用示例<a name="zh-cn_topic_0000002168254826_section14459801435"></a>
164 165 
165-- 单算子模式调用166+- 单算子模式调用
166 167 
167 ```python168 ```python
168 import os169 import os
@@ -325,7 +326,7 @@ torch_npu.npu_moe_distribute_combine(expand_x, expert_ids, expand_idx, ep_send_c
325 326 
326 ```327 ```
327 328 
328-- 图模式调用329+- 图模式调用
329 330 
330 ```python331 ```python
331 # 仅支持静态图332 # 仅支持静态图
@@ -507,4 +508,3 @@ torch_npu.npu_moe_distribute_combine(expand_x, expert_ids, expand_idx, ep_send_c
507 p.join()508 p.join()
508 print("run npu success.")509 print("run npu success.")
509 ```510 ```
510- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_distribute_combine_add_rms_norm.md+78-80
@@ -8,12 +8,12 @@
8 8 
9## 功能说明<a name="zh-cn_topic_0000002322738573_section1470016430218"></a>9## 功能说明<a name="zh-cn_topic_0000002322738573_section1470016430218"></a>
10 10 
11-- API功能:11+- API功能:
12 12 
13 需与[torch_npu.npu_moe_distribute_dispatch_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)配套使用,相当于按`npu_moe_distribute_dispatch_v2`算子收集数据的路径原路返回后对数据进行`add_rms_norm`操作。13 需与[torch_npu.npu_moe_distribute_dispatch_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)配套使用,相当于按`npu_moe_distribute_dispatch_v2`算子收集数据的路径原路返回后对数据进行`add_rms_norm`操作。
14 - 支持数据整合功能,即对moe_distribute_combine、add及rms_norm进行功能融合;14 - 支持数据整合功能,即对moe_distribute_combine、add及rms_norm进行功能融合;
15 - 支持特殊专家场景。15 - 支持特殊专家场景。
16-- 计算公式:16+- 计算公式:
17 - 数据整合功能:17 - 数据整合功能:
18 18 
19 $$19 $$
@@ -24,7 +24,7 @@
24 y = \frac{x}{RMS(x)} * gamma,\quad\text{where}RMS(x) = \sqrt{\frac{1}{H}\sum_{i=1}^{H}x_{i}^{2}+norm\_eps}\\24 y = \frac{x}{RMS(x)} * gamma,\quad\text{where}RMS(x) = \sqrt{\frac{1}{H}\sum_{i=1}^{H}x_{i}^{2}+norm\_eps}\\
25 $$25 $$
26 26 
27- - 特殊专家场景:27+ - 特殊专家场景:
28 28 
29 - 零专家场景(zero_expert_num ≠ 0):29 - 零专家场景(zero_expert_num ≠ 0):
30 30 
@@ -40,128 +40,127 @@
40 40 
41## 函数原型<a name="zh-cn_topic_0000002322738573_section470115437220"></a>41## 函数原型<a name="zh-cn_topic_0000002322738573_section470115437220"></a>
42 42 
43-```43+```python
44torch_npu.npu_moe_distribute_combine_add_rms_norm(expand_x, expert_ids, expand_idx, ep_send_counts, expert_scales, residual_x, gamma, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, tp_send_counts=None, x_active_mask=None, activation_scale=None, weight_scale=None, group_list=None, expand_scales=None, shared_expert_x=None, elastic_info=None, ori_x=None, const_expert_alpha_1=None, const_expert_alpha_2=None, const_expert_v=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, global_bs=0, out_dtype=0, comm_quant_mode=0, group_list_type=0, norm_eps=1e-06, int zero_expert_num=0, int copy_expert_num=0, int const_expert_num=0) -> (Tensor, Tensor, Tensor)44torch_npu.npu_moe_distribute_combine_add_rms_norm(expand_x, expert_ids, expand_idx, ep_send_counts, expert_scales, residual_x, gamma, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, tp_send_counts=None, x_active_mask=None, activation_scale=None, weight_scale=None, group_list=None, expand_scales=None, shared_expert_x=None, elastic_info=None, ori_x=None, const_expert_alpha_1=None, const_expert_alpha_2=None, const_expert_v=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, global_bs=0, out_dtype=0, comm_quant_mode=0, group_list_type=0, norm_eps=1e-06, int zero_expert_num=0, int copy_expert_num=0, int const_expert_num=0) -> (Tensor, Tensor, Tensor)
45```45```
46 46 
47## 参数说明<a name="zh-cn_topic_0000002322738573_section187018431529"></a>47## 参数说明<a name="zh-cn_topic_0000002322738573_section187018431529"></a>
48 48 
49-- **expand\_x**(`Tensor`):必选参数,根据`expert_ids`进行扩展过的token特征,要求为2D的Tensor,shape为\(max\(`tp_world_size`, 1\) \*A, H\),数据类型支持`bfloat16`,数据格式为ND,支持非连续的Tensor。49+- **expand\_x**(`Tensor`):必选参数,根据`expert_ids`进行扩展过的token特征,要求为2D的Tensor,shape为\(max\(`tp_world_size`, 1\) \*A, H\),数据类型支持`bfloat16`,数据格式为ND,支持非连续的Tensor。
50-- **expert\_ids**(`Tensor`):必选参数,每个token的topK个专家索引,要求为2D的Tensor,shape为\(BS, K\)。数据类型支持`int32`,数据格式为ND,支持非连续的Tensor。对应`torch_npu.npu_moe_distribute_dispatch`的`expert_ids`输入,张量里value取值范围为\[0, `moe_expert_num`\),且同一行中的K个value不能重复。50+- **expert\_ids**(`Tensor`):必选参数,每个token的topK个专家索引,要求为2D的Tensor,shape为\(BS, K\)。数据类型支持`int32`,数据格式为ND,支持非连续的Tensor。对应`torch_npu.npu_moe_distribute_dispatch`的`expert_ids`输入,张量里value取值范围为\[0, `moe_expert_num`\),且同一行中的K个value不能重复。
51-- **expand\_idx**(`Tensor`):必选参数,表示给同一专家发送的token个数,要求是1D的Tensor,shape为\(A\*128, \)。数据类型支持int32,数据格式为ND,支持非连续的Tensor。对应`torch_npu.npu_moe_distribute_dispatch`的`expand_idx`输出。51+- **expand\_idx**(`Tensor`):必选参数,表示给同一专家发送的token个数,要求是1D的Tensor,shape为\(A\*128, \)。数据类型支持int32,数据格式为ND,支持非连续的Tensor。对应`torch_npu.npu_moe_distribute_dispatch`的`expand_idx`输出。
52-- **ep\_send\_counts**(`Tensor`):必选参数,表示本卡每个专家发给EP(Expert Parallelism)域每个卡的数据量,要求是1D的Tensor 。数据类型支持`int32`,数据格式为ND,支持非连续的Tensor。对应`torch_npu.npu_moe_distribute_dispatch`的`ep_recv_counts`输出。52+- **ep\_send\_counts**(`Tensor`):必选参数,表示本卡每个专家发给EP(Expert Parallelism)域每个卡的数据量,要求是1D的Tensor 。数据类型支持`int32`,数据格式为ND,支持非连续的Tensor。对应`torch_npu.npu_moe_distribute_dispatch`的`ep_recv_counts`输出。
53 <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(`ep_world_size`\*max\(`tp_world_size`, 1\)\*local\_expert\_num, \)。53 <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(`ep_world_size`\*max\(`tp_world_size`, 1\)\*local\_expert\_num, \)。
54 54 
55-- **expert\_scales**(`Tensor`):必选参数,表示每个token的topK个专家的权重,要求是2D的Tensor,shape为\(BS, K\),其中共享专家不需要乘权重系数,直接相加即可。数据类型支持`float`,数据格式为ND,支持非连续的Tensor。55+- **expert\_scales**(`Tensor`):必选参数,表示每个token的topK个专家的权重,要求是2D的Tensor,shape为\(BS, K\),其中共享专家不需要乘权重系数,直接相加即可。数据类型支持`float`,数据格式为ND,支持非连续的Tensor。
56-- **residual\_x**(`Tensor`):必选参数,表示处理后的token需要add的参数,要求是3D的Tensor,shape为\(BS, 1, H\)。数据类型支持`bfloat16`,数据格式为ND,支持非连续的Tensor。56+- **residual\_x**(`Tensor`):必选参数,表示处理后的token需要add的参数,要求是3D的Tensor,shape为\(BS, 1, H\)。数据类型支持`bfloat16`,数据格式为ND,支持非连续的Tensor。
57-- **gamma**(`Tensor`):必选参数,表示rms\_norm的权重,要求是1D的Tensor,shape为\(H, \)。数据类型支持`bfloat16`,数据格式为ND,支持非连续的Tensor。57+- **gamma**(`Tensor`):必选参数,表示rms\_norm的权重,要求是1D的Tensor,shape为\(H, \)。数据类型支持`bfloat16`,数据格式为ND,支持非连续的Tensor。
58-- **group\_ep**(`str`):必选参数,EP通信域名称,专家并行的通信域。字符串长度范围为\[1, 128\),不能和`group_tp`相同。58+- **group\_ep**(`str`):必选参数,EP通信域名称,专家并行的通信域。字符串长度范围为\[1, 128\),不能和`group_tp`相同。
59-- **ep\_world\_size**(`int`):必选参数,EP通信域size。59+- **ep\_world\_size**(`int`):必选参数,EP通信域size。
60- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持\[2, 768\]。60+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持\[2, 768\]。
61 61 
62-- **ep\_rank\_id**(`int`):必选参数,EP通信域本卡ID,取值范围\[0, `ep_world_size`\),同一个EP通信域中各卡的ep\_rank\_id不重复。62+- **ep\_rank\_id**(`int`):必选参数,EP通信域本卡ID,取值范围\[0, `ep_world_size`\),同一个EP通信域中各卡的ep\_rank\_id不重复。
63-- **moe\_expert\_num**(`int`):必选参数,MoE专家数量,取值范围\[1, 1024\],并且满足`moe_expert_num`%\(`ep_world_size`-`shared_expert_rank_num`\)=0。63+- **moe\_expert\_num**(`int`):必选参数,MoE专家数量,取值范围\[1, 1024\],并且满足`moe_expert_num`%\(`ep_world_size`-`shared_expert_rank_num`\)=0。
64-- **tp\_send\_counts**(`Tensor`):可选参数,表示本卡每个专家发给TP(Tensor Parallelism)通信域每个卡的数据量。对应`torch_npu.npu_moe_distribute_dispatch`的`tp_recv_counts`输出。64+- **tp\_send\_counts**(`Tensor`):可选参数,表示本卡每个专家发给TP(Tensor Parallelism)通信域每个卡的数据量。对应`torch_npu.npu_moe_distribute_dispatch`的`tp_recv_counts`输出。
65- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求是一个1D Tensor,shape为\(`tp_world_size`, \),数据类型支持`int32`,数据格式要求为ND,支持非连续的Tensor。65+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求是一个1D Tensor,shape为\(`tp_world_size`, \),数据类型支持`int32`,数据格式要求为ND,支持非连续的Tensor。
66 66 
67-- **x\_active\_mask**(`Tensor`):Tensor类型,67+- **x\_active\_mask**(`Tensor`):Tensor类型,
68- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求是一个1D或者2D Tensor。当输入为1D时,shape为\(BS, \); 当输入为2D时,shape为\(BS, K\)。数据类型支持bool,数据格式要求为ND,支持非连续的Tensor。当输入为1D时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2D时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。68+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求是一个1D或者2D Tensor。当输入为1D时,shape为\(BS, \); 当输入为2D时,shape为\(BS, K\)。数据类型支持bool,数据格式要求为ND,支持非连续的Tensor。当输入为1D时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2D时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。
69 69 
70-- **activation\_scale**(`Tensor`):可选参数,**预留参数暂未使用,使用默认值即可。**70+- **activation\_scale**(`Tensor`):可选参数,**预留参数暂未使用,使用默认值即可。**
71-- **weight\_scale**(`Tensor`):可选参数,**预留参数暂未使用,使用默认值即可。**71+- **weight\_scale**(`Tensor`):可选参数,**预留参数暂未使用,使用默认值即可。**
72-- **group\_list**(`Tensor`):可选参数,**预留参数暂未使用,使用默认值即可。**72+- **group\_list**(`Tensor`):可选参数,**预留参数暂未使用,使用默认值即可。**
73-- **expand\_scales**(`Tensor`):可选参数,对应`torch_npu.npu_moe_distribute_dispatch`的`expand_scales`输出。73+- **expand\_scales**(`Tensor`):可选参数,对应`torch_npu.npu_moe_distribute_dispatch`的`expand_scales`输出。
74- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:暂不支持该参数,使用默认值即可。74+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:暂不支持该参数,使用默认值即可。
75 75 
76-- **shared\_expert\_x**(`Tensor`):可选参数,数据类型需与`expand_x`保持一致。仅在共享专家卡数量`shared_expert_rank_num`为0的场景下使用,表示共享专家token,在combine时需要加上。76+- **shared\_expert\_x**(`Tensor`):可选参数,数据类型需与`expand_x`保持一致。仅在共享专家卡数量`shared_expert_rank_num`为0的场景下使用,表示共享专家token,在combine时需要加上。
77- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型需与`expand_x`保持一致,shape为\[BS, H\]。77+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型需与`expand_x`保持一致,shape为\[BS, H\]。
78 78 
79-- **elastic\_info** (`Tensor`):预留参数,当前版本不支持,传默认值None即可。79+- **elastic\_info** (`Tensor`):预留参数,当前版本不支持,传默认值None即可。
80 80 
81-- **ori\_x** (`Tensor`):可选参数,表示未经过FFN的token数据,在使能copy_expert或使能const_expert的场景下需要本输入数据。可选择传入有效数据或填空指针,当copy_expert_num不为零或const_expert_num不为零时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(BS,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。81+- **ori\_x** (`Tensor`):可选参数,表示未经过FFN的token数据,在使能copy_expert或使能const_expert的场景下需要本输入数据。可选择传入有效数据或填空指针,当copy_expert_num不为零或const_expert_num不为零时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(BS,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。
82 82 
83-- **const\_expert\_alpha\_1** (`Tensor`):可选参数,在使能const_expert的场景下需要输入的计算系数。可选择传入有效数据或填None,当const_expert_num不为零时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。83+- **const\_expert\_alpha\_1** (`Tensor`):可选参数,在使能const_expert的场景下需要输入的计算系数。可选择传入有效数据或填None,当const_expert_num不为零时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。
84 84 
85-- **const\_expert\_alpha\_2** (`Tensor`):可选参数,在使能const_expert的场景下需要输入的计算系数。可选择传入有效数据或填None,当const_expert_num不为零时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。85+- **const\_expert\_alpha\_2** (`Tensor`):可选参数,在使能const_expert的场景下需要输入的计算系数。可选择传入有效数据或填None,当const_expert_num不为零时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。
86 86 
87-- **const\_expert\_v** (`Tensor`):可选参数,在使能const_expert的场景下需要输入的计算系数。可选择传入有效数据或填None,当const_expert_num不为零时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。87+- **const\_expert\_v** (`Tensor`):可选参数,在使能const_expert的场景下需要输入的计算系数。可选择传入有效数据或填None,当const_expert_num不为零时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。
88 88 
89-- **group\_tp**(`str`):可选参数,TP通信域名称,数据并行的通信域。有TP域通信才需要传参。89+- **group\_tp**(`str`):可选参数,TP通信域名称,数据并行的通信域。有TP域通信才需要传参。
90- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。90+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。
91 91 
92-- **tp\_world\_size**(`int`):可选参数,TP通信域size。有TP域通信才需要传参。92+- **tp\_world\_size**(`int`):可选参数,TP通信域size。有TP域通信才需要传参。
93- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。93+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。
94 94 
95-- **tp\_rank\_id**(`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。95+- **tp\_rank\_id**(`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。
96- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],同一个TP通信域中各卡的tp\_rank\_id不重复。无TP域通信时,传0即可。96+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],同一个TP通信域中各卡的tp\_rank\_id不重复。无TP域通信时,传0即可。
97 97 
98-- **expert\_shard\_type**(`int`):可选参数,表示共享专家卡排布类型。98+- **expert\_shard\_type**(`int`):可选参数,表示共享专家卡排布类型。
99- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。99+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。
100 100 
101-- **shared\_expert\_num**(`int`):可选参数,表示共享专家数量,一个共享专家可以复制部署到多个卡上。**预留参数暂未使用,仅支持默认值0。**101+- **shared\_expert\_num**(`int`):可选参数,表示共享专家数量,一个共享专家可以复制部署到多个卡上。**预留参数暂未使用,仅支持默认值0。**
102-- **shared\_expert\_rank\_num**(`int`):可选参数,表示共享专家卡数量。**预留参数暂未使用,仅支持默认值0。**102+- **shared\_expert\_rank\_num**(`int`):可选参数,表示共享专家卡数量。**预留参数暂未使用,仅支持默认值0。**
103-- **global\_bs**(`int`):可选参数,表示EP域全局的batch size大小。103+- **global\_bs**(`int`):可选参数,表示EP域全局的batch size大小。
104- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*`ep_world_size`,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*`ep_world_size`。104+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*`ep_world_size`,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*`ep_world_size`。
105 105 
106-- **out\_dtype**(`int`):可选参数,**预留参数暂未使用,使用默认值即可**。106+- **out\_dtype**(`int`):可选参数,**预留参数暂未使用,使用默认值即可**。
107-- **comm\_quant\_mode**(`int`):可选参数,表示通信量化类型。**预留参数暂未使用,使用默认值即可**。107+- **comm\_quant\_mode**(`int`):可选参数,表示通信量化类型。**预留参数暂未使用,使用默认值即可**。
108-- **group\_list\_type**(`int`):可选参数,**预留参数暂未使用,使用默认值即可**。108+- **group\_list\_type**(`int`):可选参数,**预留参数暂未使用,使用默认值即可**。
109-- **norm\_eps**(`float`):可选参数,用于防止add\_rms\_norm除0错误,默认值为1e-6。109+- **norm\_eps**(`float`):可选参数,用于防止add\_rms\_norm除0错误,默认值为1e-6。
110 110 
111-- **zero\_expert\_num** (`int`):可选参数,表示零专家的数量。取值范围\[0, MAX_INT32\),其中MAX_INT32值为2147483647,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。111+- **zero\_expert\_num** (`int`):可选参数,表示零专家的数量。取值范围\[0, MAX_INT32\),其中MAX_INT32值为2147483647,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。
112 112 
113-- **copy\_expert\_num** (`int`):可选参数,表示copy专家的数量。取值范围\[0, MAX_INT32\),其中MAX_INT32值为2147483647,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。113+- **copy\_expert\_num** (`int`):可选参数,表示copy专家的数量。取值范围\[0, MAX_INT32\),其中MAX_INT32值为2147483647,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。
114 114 
115-- **const\_expert\_num** (`int`):可选参数,表示常量专家的数量。取值范围\[0, MAX_INT32\),其中MAX_INT32值为2147483647,合法的常量专家的ID值是\[moe\_expert\_num+zero\_expert\_num+copy\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num+const\_expert\_num\)。115+- **const\_expert\_num** (`int`):可选参数,表示常量专家的数量。取值范围\[0, MAX_INT32\),其中MAX_INT32值为2147483647,合法的常量专家的ID值是\[moe\_expert\_num+zero\_expert\_num+copy\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num+const\_expert\_num\)。
116 116 
117## 返回值说明<a name="zh-cn_topic_0000002322738573_section1370204314220"></a>117## 返回值说明<a name="zh-cn_topic_0000002322738573_section1370204314220"></a>
118 118 
119-- **y**(`Tensor`):表示combine处理后的token进行add\_rms\_norm计算后的结果,要求是3D的Tensor,shape为\(BS, 1, H\),数据类型与输入`residual_x`保持一致,数据格式为ND,不支持非连续的Tensor。119+- **y**(`Tensor`):表示combine处理后的token进行add\_rms\_norm计算后的结果,要求是3D的Tensor,shape为\(BS, 1, H\),数据类型与输入`residual_x`保持一致,数据格式为ND,不支持非连续的Tensor。
120-- **rstd\_out**(`Tensor`):表示add\_rms\_norm的输出结果,要求是3D的Tensor,shape为\(BS, 1, 1\),数据类型支持`float`,数据格式为ND,不支持非连续的Tensor。120+- **rstd\_out**(`Tensor`):表示add\_rms\_norm的输出结果,要求是3D的Tensor,shape为\(BS, 1, 1\),数据类型支持`float`,数据格式为ND,不支持非连续的Tensor。
121-- **x**(`Tensor`):表示combine处理后的token进行add计算后的结果,要求是3D的Tensor,shape为\(BS, 1, H\),数据类型与输入`residual_x`保持一致,数据格式为ND,不支持非连续的Tensor。121+- **x**(`Tensor`):表示combine处理后的token进行add计算后的结果,要求是3D的Tensor,shape为\(BS, 1, H\),数据类型与输入`residual_x`保持一致,数据格式为ND,不支持非连续的Tensor。
122 122 
123## 约束说明<a name="zh-cn_topic_0000002322738573_section470214314214"></a>123## 约束说明<a name="zh-cn_topic_0000002322738573_section470214314214"></a>
124 124 
125-- 该接口支持推理场景下使用。125+- 该接口支持推理场景下使用。
126-- 该接口支持图模式。126+- 该接口支持图模式。
127- 调用接口过程中使用的expert_ids、x_active_mask、elastic_info、group_ep、ep_world_size、moe_expert_num、group_tp、tp_world_size、expert_shard_type、shared_expert_num、shared_expert_rank_num、global_bs、comm_alg、zero_expert_num、copy_expert_num、const_expert_num参数、HCCL_BUFFSIZE取值所有卡需保持一致,网络中不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)对应参数也保持一致。127- 调用接口过程中使用的expert_ids、x_active_mask、elastic_info、group_ep、ep_world_size、moe_expert_num、group_tp、tp_world_size、expert_shard_type、shared_expert_num、shared_expert_rank_num、global_bs、comm_alg、zero_expert_num、copy_expert_num、const_expert_num参数、HCCL_BUFFSIZE取值所有卡需保持一致,网络中不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)对应参数也保持一致。
128-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。128+- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。
129-- moe_expert_num + zero_expert_num + copy_expert_num + const_expert_num < MAX_INT32,其中MAX_INT32值为2147483647。129+- moe_expert_num + zero_expert_num + copy_expert_num + const_expert_num < MAX_INT32,其中MAX_INT32值为2147483647。
130-- 参数里Shape使用的变量如下:130+- 参数里Shape使用的变量如下:
131 - A:表示本卡需要分发的最大token数量,取值范围如下:131 - A:表示本卡需要分发的最大token数量,取值范围如下:
132- - 当`global_bs`为0时,要满足A >= Bs * epWorldSize * min(localExpertNum, K);132+ - 当`global_bs`为0时,要满足A >= Bs \* epWorldSize \* min(localExpertNum, K);
133 -`global_bs`非0时,要满足A >= globalBs * min(localExpertNum, K)。 133 -`global_bs`非0时,要满足A >= globalBs * min(localExpertNum, K)。
134 134 
135- - H:表示hidden size隐藏层大小。135+ - H:表示hidden size隐藏层大小。
136- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[1024, 8192\]。136+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[1024, 8192\]。
137 137 
138- - BS:表示待发送的token数量。138+ - BS:表示待发送的token数量。
139- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为0<BS≤512。139+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为0<BS≤512。
140 140 
141- - K:表示选取topK个专家,取值范围为0<K≤8同时满足0 < K ≤ moe\_expert\_num + zero_expert_num + copy_expert_num + const_expert_num。141+ - K:表示选取topK个专家,取值范围为0<K≤8同时满足0 < K ≤ moe\_expert\_num + zero_expert_num + copy_expert_num + const_expert_num。
142- - local\_expert\_num:表示本卡专家数量。142+ - local\_expert\_num:表示本卡专家数量。
143- - 对于共享专家卡,local\_expert\_num=1143+ - 对于共享专家卡,local\_expert\_num=1
144- - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num\),当local\_expert\_num\>1时,不支持TP域通信。144+ - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num\),当local\_expert\_num\>1时,不支持TP域通信。
145 145 
146-- HCCL通信域缓存区大小:146+- HCCL通信域缓存区大小:
147 147 
148 调用本接口前需检查`HCCL_BUFFSIZE`环境变量取值是否合理,该环境变量表示单个通信域占用内存大小,单位MB,不配置时默认为200MB。该场景通信域缓存区大小支持通过环境变量HCCL\_BUFFSIZE配置,也支持通过hccl_buffer_size配置(参考[《PyTorch训练模型迁移调优》](https://hiascend.com/document/redirect/canncommercial-ptmigr)中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。148 调用本接口前需检查`HCCL_BUFFSIZE`环境变量取值是否合理,该环境变量表示单个通信域占用内存大小,单位MB,不配置时默认为200MB。该场景通信域缓存区大小支持通过环境变量HCCL\_BUFFSIZE配置,也支持通过hccl_buffer_size配置(参考[《PyTorch训练模型迁移调优》](https://hiascend.com/document/redirect/canncommercial-ptmigr)中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。
149- - ep通信域内:设置大小要求 \>= 2且满足\>= 2 \* \(local\_expert\_num \* max\_bs \* ep\_world\_size \* Align512\(Align32\(2 \* H\) + 64\) + \(K + shared\_expert\_num\) \* max\_bs \* Align512\(2 \* H\)\),local\_expert\_num表示需使用MoE专家卡的本卡专家数。149+ - ep通信域内:设置大小要求 \>= 2且满足\>= 2 \* \(local\_expert\_num \* max\_bs \* ep\_world\_size \* Align512\(Align32\(2 \* H\) + 64\) + \(K + shared\_expert\_num\) \* max\_bs \* Align512\(2 \* H\)\),local\_expert\_num表示需使用MoE专家卡的本卡专家数。
150- - tp通信域内:设置大小要求 \>= (A \* Align512(Align32(h \* 2) + 44) + A \* Align512(h \* 2)) \* 2。150+ - tp通信域内:设置大小要求 \>= (A \* Align512(Align32(h \* 2) + 44) + A \* Align512(h \* 2)) \* 2。
151- - 其中 480Align512(x) = ((x+480-1)/480)\*512,Align512(x) = ((x+512-1)/512)\*512,Align32(x) = ((x+32-1)/32)\*32。151+ - 其中 480Align512(x) = ((x+480-1)/480)\*512,Align512(x) = ((x+512-1)/512)\*512,Align32(x) = ((x+32-1)/32)\*32。
152 152 
153-- 通信域使用约束:153+- 通信域使用约束:
154 154 
155- - 一个模型中的npu\_moe\_distribute\_dispatch\_v2和npu\_moe\_distribute\_combine\_add\_rms\_norm算子仅支持相同EP通信域,且该通信域中不允许有其他算子。155+ - 一个模型中的npu\_moe\_distribute\_dispatch\_v2和npu\_moe\_distribute\_combine\_add\_rms\_norm算子仅支持相同EP通信域,且该通信域中不允许有其他算子。
156 156 
157- - 一个模型中的npu\_moe\_distribute\_dispatch\_v2和npu\_moe\_distribute\_combine\_add\_rms\_norm算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。157+ - 一个模型中的npu\_moe\_distribute\_dispatch\_v2和npu\_moe\_distribute\_combine\_add\_rms\_norm算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。
158- 
159- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:一个通信域内的节点需在一个超节点内,不支持跨超节点。
160 158 
159+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:一个通信域内的节点需在一个超节点内,不支持跨超节点。
161 160 
162## 调用示例<a name="zh-cn_topic_0000002322738573_section9702174311218"></a>161## 调用示例<a name="zh-cn_topic_0000002322738573_section9702174311218"></a>
163 162 
164-- 单算子模式调用163+- 单算子模式调用
165 164 
166 ```python165 ```python
167 import os166 import os
@@ -422,7 +421,7 @@ torch_npu.npu_moe_distribute_combine_add_rms_norm(expand_x, expert_ids, expand_i
422 print("run npu success.")421 print("run npu success.")
423 ```422 ```
424 423 
425-- 图模式调用424+- 图模式调用
426 425 
427 ```python426 ```python
428 # 仅支持静态图427 # 仅支持静态图
@@ -702,4 +701,3 @@ torch_npu.npu_moe_distribute_combine_add_rms_norm(expand_x, expert_ids, expand_i
702 p.join()701 p.join()
703 print("run npu success.")702 print("run npu success.")
704 ```703 ```
705- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_distribute_combine_v2.md+115-116
@@ -9,12 +9,12 @@
9 9 
10## 功能说明<a name="zh-cn_topic_0000002168254826_section14441124184110"></a>10## 功能说明<a name="zh-cn_topic_0000002168254826_section14441124184110"></a>
11 11 
12-- API功能:12+- API功能:
13 13 
14 需与[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)配套使用,相当于按npu\_moe\_distribute\_dispatch\_v2算子收集数据的路径原路返回。14 需与[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)配套使用,相当于按npu\_moe\_distribute\_dispatch\_v2算子收集数据的路径原路返回。
15 - 支持数据整合功能,先进行reduce\_scatterv通信,再进行alltoallv通信,最后将接收的数据整合(乘权重再相加);15 - 支持数据整合功能,先进行reduce\_scatterv通信,再进行alltoallv通信,最后将接收的数据整合(乘权重再相加);
16 - 支持特殊专家场景。16 - 支持特殊专家场景。
17-- 计算公式:17+- 计算公式:
18 - 数据整合功能:18 - 数据整合功能:
19 19 
20 $rs\_out = ReduceScatterV(expand\_x)$20 $rs\_out = ReduceScatterV(expand\_x)$
@@ -37,208 +37,207 @@
37 37 
38 $Moe(ori\_x)=const\_expert\_alpha\_1*ori\_x+const\_expert\_alpha\_2*const\_expert\_v$38 $Moe(ori\_x)=const\_expert\_alpha\_1*ori\_x+const\_expert\_alpha\_2*const\_expert\_v$
39 39 
40- 
41- 
42## 函数原型<a name="zh-cn_topic_0000002168254826_section45077510411"></a>40## 函数原型<a name="zh-cn_topic_0000002168254826_section45077510411"></a>
43 41 
44-```42+```python
45torch_npu.npu_moe_distribute_combine_v2(expand_x, expert_ids, assist_info_for_combine, ep_send_counts, expert_scales, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, tp_send_counts=None, x_active_mask=None, expand_scales=None, shared_expert_x=None, elastic_info=None, ori_x=None, const_expert_alpha_1=None, const_expert_alpha_2=None, const_expert_v=None, performance_info=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, global_bs=0, comm_quant_mode=0, comm_alg="", zero_expert_num=0, copy_expert_num=0, const_expert_num=0) -> Tensor43torch_npu.npu_moe_distribute_combine_v2(expand_x, expert_ids, assist_info_for_combine, ep_send_counts, expert_scales, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, tp_send_counts=None, x_active_mask=None, expand_scales=None, shared_expert_x=None, elastic_info=None, ori_x=None, const_expert_alpha_1=None, const_expert_alpha_2=None, const_expert_v=None, performance_info=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, global_bs=0, comm_quant_mode=0, comm_alg="", zero_expert_num=0, copy_expert_num=0, const_expert_num=0) -> Tensor
46```44```
47 45 
48## 参数说明<a name="zh-cn_topic_0000002168254826_section112637109429"></a>46## 参数说明<a name="zh-cn_topic_0000002168254826_section112637109429"></a>
49 47 
50-- **expand\_x** (`Tensor`):必选参数,根据`expert_ids`进行扩展过的token特征,要求为2维张量,shape为\(max\(tp\_world\_size, 1\) \*A, H\),数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。48+- **expand\_x** (`Tensor`):必选参数,根据`expert_ids`进行扩展过的token特征,要求为2维张量,shape为\(max\(tp\_world\_size, 1\) \*A, H\),数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。
51- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家场景。49+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家场景。
52 50 
53-- **expert\_ids** (`Tensor`):必选参数,每个token的topK个专家索引,要求为2维张量,shape为\(BS, K\)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`expert_ids`输入,张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。51+- **expert\_ids** (`Tensor`):必选参数,每个token的topK个专家索引,要求为2维张量,shape为\(BS, K\)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`expert_ids`输入,张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。
54-- **assist\_info\_for\_combine** (`Tensor`):必选参数,表示给同一专家发送的token个数,要求为1维张量,shape为\(A \* 128, \)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`assist_info_for_combine`输出。52+- **assist\_info\_for\_combine** (`Tensor`):必选参数,表示给同一专家发送的token个数,要求为1维张量,shape为\(A \* 128, \)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`assist_info_for_combine`输出。
55 53 
56-- **ep\_send\_counts** (`Tensor`):必选参数,表示本卡每个专家发给EP(Expert Parallelism)域每个卡的token数(token数以前缀和的形式表示),要求为1维张量。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`ep_recv_counts`输出。54+- **ep\_send\_counts** (`Tensor`):必选参数,表示本卡每个专家发给EP(Expert Parallelism)域每个卡的token数(token数以前缀和的形式表示),要求为1维张量。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`ep_recv_counts`输出。
57- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(moe\_expert\_num+2\*global\_bs\*K\*server\_num, \),前`moe_expert_num`个数表示在EP通信域内,该卡上每个专家收到来自其他各卡的token数(以前缀和的形式表示),2\*global\_bs\*K\*server\_num用于存储机间和机内通信前,combine可提前做reduce操作的token个数和通信区偏移量,`global_bs`传入0时此处按照bs\*ep\_world\_size计算。55+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(moe\_expert\_num+2\*global\_bs\*K\*server\_num, \),前`moe_expert_num`个数表示在EP通信域内,该卡上每个专家收到来自其他各卡的token数(以前缀和的形式表示),2\*global\_bs\*K\*server\_num用于存储机间和机内通信前,combine可提前做reduce操作的token个数和通信区偏移量,`global_bs`传入0时此处按照bs\*ep\_world\_size计算。
58- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(ep\_world\_size\*max\(tp\_world\_size, 1\)\*local\_expert\_num, \)。56+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(ep\_world\_size\*max\(tp\_world\_size, 1\)\*local\_expert\_num, \)。
59 57 
60-- **expert\_scales** (`Tensor`):必选参数,表示每个token的topK个专家的权重,要求为2维张量,shape为\(BS, K\),其中共享专家不需要乘权重系数,直接相加即可。数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。58+- **expert\_scales** (`Tensor`):必选参数,表示每个token的topK个专家的权重,要求为2维张量,shape为\(BS, K\),其中共享专家不需要乘权重系数,直接相加即可。数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。
61-- **group\_ep** (`str`):必选参数,EP通信域名称,专家并行的通信域。字符串长度范围为\[1, 128\)。<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>时不能和`group_tp`相同。59+- **group\_ep** (`str`):必选参数,EP通信域名称,专家并行的通信域。字符串长度范围为\[1, 128\)。<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>时不能和`group_tp`相同。
62-- **ep\_world\_size** (`int`):必选参数,EP通信域size。60+- **ep\_world\_size** (`int`):必选参数,EP通信域size。
63- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`ep_world_size`的取值范围如下所示。61+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`ep_world_size`的取值范围如下所示。
64 - `comm_alg`设置为"fullmesh"时,`ep_world_size`取值范围为2、3、4、5、6、7、8、16、32、64、128、256。62 - `comm_alg`设置为"fullmesh"时,`ep_world_size`取值范围为2、3、4、5、6、7、8、16、32、64、128、256。
65 - `comm_alg`设置为"hierarchy"时,`ep_world_size`取值范围为16、32、64。63 - `comm_alg`设置为"hierarchy"时,`ep_world_size`取值范围为16、32、64。
66- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持\[2, 768\]。`comm_alg`设置为"hierarchy"时,取值范围为[16, 256],且为16的整数倍。64+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持\[2, 768\]。`comm_alg`设置为"hierarchy"时,取值范围为[16, 256],且为16的整数倍。
67 65 
68-- **ep\_rank\_id** (`int`):必选参数,EP通信域本卡ID,取值范围\[0, ep\_world\_size\),同一个EP通信域中各卡的`ep_rank_id`不重复。66+- **ep\_rank\_id** (`int`):必选参数,EP通信域本卡ID,取值范围\[0, ep\_world\_size\),同一个EP通信域中各卡的`ep_rank_id`不重复。
69-- **moe\_expert\_num** (`int`):必选参数,MoE专家数量,取值范围\[1, 1024\],并且满足以下条件:moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0。67+- **moe\_expert\_num** (`int`):必选参数,MoE专家数量,取值范围\[1, 1024\],并且满足以下条件:moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0。
70- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`comm_alg`设置为"hierarchy"时,取值范围为(0, 512]。68+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`comm_alg`设置为"hierarchy"时,取值范围为(0, 512]。
71- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。69- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
72-- **tp\_send\_counts** (`Tensor`):可选参数,表示本卡每个专家发给TP(Tensor Parallelism)通信域每个卡的数据量。对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`tp_recv_counts`输出。70+- **tp\_send\_counts** (`Tensor`):可选参数,表示本卡每个专家发给TP(Tensor Parallelism)通信域每个卡的数据量。对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`tp_recv_counts`输出。
73- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP通信域,使用默认输入。71+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP通信域,使用默认输入。
74- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求为一个1维张量,shape为\(tp\_world\_size, \),数据类型支持`int32`,数据格式要求为$ND$,支持非连续的Tensor。72+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求为一个1维张量,shape为\(tp\_world\_size, \),数据类型支持`int32`,数据格式要求为$ND$,支持非连续的Tensor。
75 73 
76-- **x\_active\_mask** (`Tensor`):可选参数,表示token是否参与通信。74+- **x\_active\_mask** (`Tensor`):可选参数,表示token是否参与通信。
77- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:75+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
78 - `comm_alg`设置为"fullmesh"时,要求为一个1维或2维张量。当输入为1维时,shape为\(BS, \); 当输入为2维时,shape为\(BS, K\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。当输入为1维时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2维时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。支持2维张量属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。76 - `comm_alg`设置为"fullmesh"时,要求为一个1维或2维张量。当输入为1维时,shape为\(BS, \); 当输入为2维时,shape为\(BS, K\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。当输入为1维时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2维时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。支持2维张量属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。
79 - `comm_alg`设置为"hierarchy"时,当前版本不支持,使用默认值None即可。77 - `comm_alg`设置为"hierarchy"时,当前版本不支持,使用默认值None即可。
80- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求为一个1维或2维张量。当输入为1维时,shape为\(BS, \); 当输入为2维时,shape为\(BS, K\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。当输入为1维时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2维时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。78+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求为一个1维或2维张量。当输入为1维时,shape为\(BS, \); 当输入为2维时,shape为\(BS, K\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。当输入为1维时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2维时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。
81 79 
82-- **expand\_scales** (`Tensor`):可选参数,对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`expand_scales`输出。80+- **expand\_scales** (`Tensor`):可选参数,对应[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)的`expand_scales`输出。
83- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:必选参数,要求为1维张量,shape为\(A, \),数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。81+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:必选参数,要求为1维张量,shape为\(A, \),数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。
84- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`comm_alg`设置为"hierarchy"时,要求为1维张量,shape为\(A, \),数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。`comm_alg`设置为""时,暂不支持该参数,使用默认值即可。82+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`comm_alg`设置为"hierarchy"时,要求为1维张量,shape为\(A, \),数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。`comm_alg`设置为""时,暂不支持该参数,使用默认值即可。
85 83 
86-- **shared\_expert\_x** (`Tensor`):可选参数,数据类型需与`expand_x`保持一致。仅在共享专家卡数量`shared_expert_rank_num`为0的场景下使用,表示共享专家token,在combine\_v2后需要做add的值。84+- **shared\_expert\_x** (`Tensor`):可选参数,数据类型需与`expand_x`保持一致。仅在共享专家卡数量`shared_expert_rank_num`为0的场景下使用,表示共享专家token,在combine\_v2后需要做add的值。
87- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。85+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。
88- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求为一个2维或3维的张量,当张量为2D时,shape为\(BS, H\);当张量为3D时,前两位的乘积需等于BS,第三维需等于H。86+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求为一个2维或3维的张量,当张量为2D时,shape为\(BS, H\);当张量为3D时,前两位的乘积需等于BS,第三维需等于H。
89 87 
90-- **elastic\_info** (`Tensor`):预留参数,当前版本不支持,传默认值None即可。88+- **elastic\_info** (`Tensor`):预留参数,当前版本不支持,传默认值None即可。
91 89 
92-- **ori\_x** (`Tensor`):可选参数,表示未经过FFN的token数据,在`copy_expert_num`不为0或`const_expert_num`不为0的场景下需要本输入数据。90+- **ori\_x** (`Tensor`):可选参数,表示未经过FFN的token数据,在`copy_expert_num`不为0或`const_expert_num`不为0的场景下需要本输入数据。
93- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:91+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
94 - `comm_alg`设置为"fullmesh"时,可选择传入有效数据或填None,当`copy_expert_num`不为0时必须传入有效数据;当传入有效数据时,要求是一个2D的Tensor,shape为(BS,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。参数为有效数据时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。92 - `comm_alg`设置为"fullmesh"时,可选择传入有效数据或填None,当`copy_expert_num`不为0时必须传入有效数据;当传入有效数据时,要求是一个2D的Tensor,shape为(BS,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。参数为有效数据时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。
95 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值None即可。93 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值None即可。
96- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:可选择传入有效数据或填None,当`copy_expert_num`不为0或`const_expert_num`不为0时必须传入有效数据;当传入有效数据时,要求是一个2D的Tensor,shape为(BS,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。94+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:可选择传入有效数据或填None,当`copy_expert_num`不为0或`const_expert_num`不为0时必须传入有效数据;当传入有效数据时,要求是一个2D的Tensor,shape为(BS,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。
97 95 
98-- **const\_expert\_alpha\_1** (`Tensor`):可选参数,在`const_expert_num`不为0的场景下需要输入的计算系数。 96+- **const\_expert\_alpha\_1** (`Tensor`):可选参数,在`const_expert_num`不为0的场景下需要输入的计算系数。
99- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:预留参数,当前版本不支持,传None即可。97+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:预留参数,当前版本不支持,传None即可。
100- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:可选择传入有效数据或填None,当`const_expert_num`不为0时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。98+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:可选择传入有效数据或填None,当`const_expert_num`不为0时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。
101 99 
102-- **const\_expert\_alpha\_2** (`Tensor`):可选参数,在`const_expert_num`不为0的场景下需要输入的计算系数。 100+- **const\_expert\_alpha\_2** (`Tensor`):可选参数,在`const_expert_num`不为0的场景下需要输入的计算系数。
103- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:预留参数,当前版本不支持,传None即可。101+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:预留参数,当前版本不支持,传None即可。
104- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:可选择传入有效数据或填None,当`const_expert_num`不为0时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。102+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:可选择传入有效数据或填None,当`const_expert_num`不为0时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。
105 103 
106-- **const\_expert\_v** (`Tensor`):可选参数,在`const_expert_num`不为0的场景下需要输入的计算系数。 104+- **const\_expert\_v** (`Tensor`):可选参数,在`const_expert_num`不为0的场景下需要输入的计算系数。
107- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:预留参数,当前版本不支持,传None即可。105+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:预留参数,当前版本不支持,传None即可。
108- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:可选择传入有效数据或填None,当`const_expert_num`不为0时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。106+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:可选择传入有效数据或填None,当`const_expert_num`不为0时必须传入有效输入;当传入有效数据时,要求是一个2D的Tensor,shape为(const_expert_num,H),数据类型需跟expand_x保持一致;数据格式要求为ND,支持非连续的Tensor。
109 107 
110-- **performance\_info** (`Tensor`):可选参数,表示本卡等待各卡数据的通信时间,单位为us(微秒)。单次算子调用各卡通信耗时会累加到该Tensor上,算子内部不进行自动清零,因此每次启用此Tensor开始记录耗时前需对Tensor清零。当传入None时表示不使能记录通信耗时功能;当传入有效数据时,要求是一个1D的Tensor,shape为(ep_world_size,),数据类型支持int64,数据格式要求为ND,支持非连续的Tensor。108+- **performance\_info** (`Tensor`):可选参数,表示本卡等待各卡数据的通信时间,单位为us(微秒)。单次算子调用各卡通信耗时会累加到该Tensor上,算子内部不进行自动清零,因此每次启用此Tensor开始记录耗时前需对Tensor清零。当传入None时表示不使能记录通信耗时功能;当传入有效数据时,要求是一个1D的Tensor,shape为(ep_world_size,),数据类型支持int64,数据格式要求为ND,支持非连续的Tensor。
111 109 
112-- **group\_tp** (`string`):可选参数,TP通信域名称,数据并行的通信域。有TP域通信才需要传参,若无TP域通信,使用默认值""即可。110+- **group\_tp** (`string`):可选参数,TP通信域名称,数据并行的通信域。有TP域通信才需要传参,若无TP域通信,使用默认值""即可。
113- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:eager模式使用默认值即可,图模式传入与`group_ep`相同。111+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:eager模式使用默认值即可,图模式传入与`group_ep`相同。
114- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。112+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。
115 113 
116-- **tp\_world\_size** (`int`):可选参数,TP通信域size。有TP域通信才需要传参。114+- **tp\_world\_size** (`int`):可选参数,TP通信域size。有TP域通信才需要传参。
117- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。115+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。
118- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。116+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。
119 117 
120-- **tp\_rank\_id** (`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。118+- **tp\_rank\_id** (`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。
121- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。119+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。
122- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],同一个TP通信域中各卡的`tp_rank_id`不重复。无TP域通信时,传0即可。120+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],同一个TP通信域中各卡的`tp_rank_id`不重复。无TP域通信时,传0即可。
123 121 
124-- **expert\_shard\_type** (`int`):可选参数,表示共享专家卡排布类型。122+- **expert\_shard\_type** (`int`):可选参数,表示共享专家卡排布类型。
125- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。123+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。
126- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。124+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。
127 125 
128-- **shared\_expert\_num** (`int`):可选参数,表示共享专家数量,一个共享专家可以复制部署到多个卡上。126+- **shared\_expert\_num** (`int`):可选参数,表示共享专家数量,一个共享专家可以复制部署到多个卡上。
129- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。127+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。
130- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, 4\],0表示无共享专家,默认值为1。128+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, 4\],0表示无共享专家,默认值为1。
131 129 
132-- **shared\_expert\_rank\_num** (`int`):可选参数,表示共享专家卡数量。130+- **shared\_expert\_rank\_num** (`int`):可选参数,表示共享专家卡数量。
133- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家,使用默认值即可。131+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家,使用默认值即可。
134- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, ep\_world\_size\)。取0表示无共享专家,不取0需满足shared\_expert\_rank\_num%shared\_expert\_num=0。132+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, ep\_world\_size\)。取0表示无共享专家,不取0需满足shared\_expert\_rank\_num%shared\_expert\_num=0。
135 133 
136-- **global\_bs** (`int`):可选参数,表示EP域全局的batch size大小。当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。134+- **global\_bs** (`int`):可选参数,表示EP域全局的batch size大小。当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。
137 135 
138-- **comm\_quant\_mode** (`int`):可选参数,表示通信量化类型。136+- **comm\_quant\_mode** (`int`):可选参数,表示通信量化类型。
139- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:支持取0和2。0表示通信时不量化,2表示通信时进行`int8`量化。仅当`comm_alg`配置为"hierarchy"或HCCL\_INTRA\_PCIE\_ENABLE=1且HCCL\_INTRA\_ROCE\_ENABLE=0且驱动版本不低于25.0.RC1.1时才支持取2。137+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:支持取0和2。0表示通信时不量化,2表示通信时进行`int8`量化。仅当`comm_alg`配置为"hierarchy"或HCCL\_INTRA\_PCIE\_ENABLE=1且HCCL\_INTRA\_ROCE\_ENABLE=0且驱动版本不低于25.0.RC1.1时才支持取2。
140- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持取0和2。0表示通信时不量化,2表示通信时进行`int8`量化。当且仅当`tp_world_size`不等于2时,可以使能`int8`量化。138+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持取0和2。0表示通信时不量化,2表示通信时进行`int8`量化。当且仅当`tp_world_size`不等于2时,可以使能`int8`量化。
141 139 
142-- **comm\_alg** (`str`):可选参数,表示通信亲和内存布局算法。140+- **comm\_alg** (`str`):可选参数,表示通信亲和内存布局算法。
143- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当前版本支持"","fullmesh","hierarchy"三种输入方式。推荐配置"hierarchy"并搭配25.0.RC1.1及以上版本驱动使用。141+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当前版本支持"","fullmesh","hierarchy"三种输入方式。推荐配置"hierarchy"并搭配25.0.RC1.1及以上版本驱动使用。
144 - "": 配置HCCL\_INTRA\_PCIE\_ENABLE=1和HCCL\_INTRA\_ROCE\_ENABLE=0时,调用"hierarchy"算法,否则调用"fullmesh"算法。不推荐使用该方式。142 - "": 配置HCCL\_INTRA\_PCIE\_ENABLE=1和HCCL\_INTRA\_ROCE\_ENABLE=0时,调用"hierarchy"算法,否则调用"fullmesh"算法。不推荐使用该方式。
145 - "fullmesh": token数据直接通过RDMA方式发回目标卡。143 - "fullmesh": token数据直接通过RDMA方式发回目标卡。
146 - "hierarchy": token数据经过机内、跨机两次发送,先在server内将同一个token数据汇总求和,再跨机发送,以减少跨机数据量。144 - "hierarchy": token数据经过机内、跨机两次发送,先在server内将同一个token数据汇总求和,再跨机发送,以减少跨机数据量。
147- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前版本支持"","hierarchy"两种输入方式。145+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前版本支持"","hierarchy"两种输入方式。
148 - "": 默认值,token数据直接通过MTE方式发回目标卡。146 - "": 默认值,token数据直接通过MTE方式发回目标卡。
149 - "hierarchy": token数据经过机内、跨机两次发送,先在server内将同一个token数据汇总求和,再跨机发送,以减少跨机数据量。模板仅支持`tp_world_size`为1、共享专家为0的场景,且不支持二维mask、特殊专家、动态缩容、性能打点场景。147 - "hierarchy": token数据经过机内、跨机两次发送,先在server内将同一个token数据汇总求和,再跨机发送,以减少跨机数据量。模板仅支持`tp_world_size`为1、共享专家为0的场景,且不支持二维mask、特殊专家、动态缩容、性能打点场景。
150 148 
151-- **zero\_expert\_num** (`int`):可选参数,表示零专家的数量。149+- **zero\_expert\_num** (`int`):可选参数,表示零专家的数量。
152- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:150+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
153 - `comm_alg`设置为"fullmesh"时,取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。参数为非0时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。151 - `comm_alg`设置为"fullmesh"时,取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。参数为非0时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。
154 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值0即可。152 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值0即可。
155- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。153+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。
156 154 
157-- **copy\_expert\_num** (`int`):可选参数,表示拷贝专家的数量。155+- **copy\_expert\_num** (`int`):可选参数,表示拷贝专家的数量。
158- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:156+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
159 - `comm_alg`设置为"fullmesh"时,取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。参数为非0时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。157 - `comm_alg`设置为"fullmesh"时,取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。参数为非0时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。
160 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值0即可。158 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值0即可。
161- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。159+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。
162 160 
163-- **const\_expert\_num** (`int`):可选参数,表示常量专家的数量。 161+- **const\_expert\_num** (`int`):可选参数,表示常量专家的数量。
164- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当前版本不支持,传0即可。162+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当前版本不支持,传0即可。
165- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的常量专家的ID值是\[moe\_expert\_num+zero\_expert\_num+copy\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num+const\_expert\_num\)。163+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的常量专家的ID值是\[moe\_expert\_num+zero\_expert\_num+copy\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num+const\_expert\_num\)。
166 164 
167## 返回值说明<a name="zh-cn_topic_0000002168254826_section22231435517"></a>165## 返回值说明<a name="zh-cn_topic_0000002168254826_section22231435517"></a>
166+ 
168`Tensor`167`Tensor`
169 168 
170表示处理后的token,要求为2维张量,shape为\(BS, H\),数据类型支持`bfloat16``float16`,类型与输入`expand_x`保持一致,数据格式为$ND$,不支持非连续的Tensor。169表示处理后的token,要求为2维张量,shape为\(BS, H\),数据类型支持`bfloat16``float16`,类型与输入`expand_x`保持一致,数据格式为$ND$,不支持非连续的Tensor。
171 170 
172## 约束说明<a name="zh-cn_topic_0000002168254826_section12345537164214"></a>171## 约束说明<a name="zh-cn_topic_0000002168254826_section12345537164214"></a>
173 172 
174-- 该接口支持推理场景下使用。173+- 该接口支持推理场景下使用。
175-- 该接口支持静态图模式,`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`必须配套使用。174+- 该接口支持静态图模式,`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`必须配套使用。
176-- 在不同产品型号、不同通信算法或不同版本中,`npu_moe_distribute_dispatch_v2`的Tensor输出`assist_info_for_combine`、`ep_recv_counts`、`tp_recv_counts`、`expand_scales`中的元素值可能不同,使用时直接将上述Tensor传给`npu_moe_distribute_combine_v2`对应参数即可,模型其他业务逻辑不应对其存在依赖。175+- 在不同产品型号、不同通信算法或不同版本中,`npu_moe_distribute_dispatch_v2`的Tensor输出`assist_info_for_combine`、`ep_recv_counts`、`tp_recv_counts`、`expand_scales`中的元素值可能不同,使用时直接将上述Tensor传给`npu_moe_distribute_combine_v2`对应参数即可,模型其他业务逻辑不应对其存在依赖。
177-- 调用接口过程中使用的`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`shared_expert_num`、`shared_expert_rank_num`、`global_bs`参数取值所有卡需保持一致,`group_ep`、`ep_world_size`、`group_tp`、`tp_world_size`、`expert_shard_type`、`global_bs`网络中不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)对应参数也保持一致。176+- 调用接口过程中使用的`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`shared_expert_num`、`shared_expert_rank_num`、`global_bs`参数取值所有卡需保持一致,`group_ep`、`ep_world_size`、`group_tp`、`tp_world_size`、`expert_shard_type`、`global_bs`网络中不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_dispatch\_v2](torch_npu-npu_moe_distribute_dispatch_v2.md)对应参数也保持一致。
178-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。177+- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。
179-- moe_expert_num + zero_expert_num + copy_expert_num + const_expert_num < MAX_INT32。178+- moe_expert_num + zero_expert_num + copy_expert_num + const_expert_num < MAX_INT32。
180-- 参数里Shape使用的变量如下:179+- 参数里Shape使用的变量如下:
181- - A:表示本卡接收的最大token数量,取值范围如下180+ - A:表示本卡接收的最大token数量,取值范围如下
182- - 对于共享专家,要满足A=BS\*ep\_world\_size*shared\_expert\_num/shared\_expert\_rank\_num。181+ - 对于共享专家,要满足A=BS\*ep\_world\_size*shared\_expert\_num/shared\_expert\_rank\_num。
183- - 对于MoE专家,当global\_bs为0时,要满足A\>=BS\*ep\_world\_size\*min\(local\_expert\_num, K\);当global\_bs不为0时,要满足A\>=global\_bs\* min\(local\_expert\_num, K\)。182+ - 对于MoE专家,当global\_bs为0时,要满足A\>=BS\*ep\_world\_size\*min\(local\_expert\_num, K\);当global\_bs不为0时,要满足A\>=global\_bs\* min\(local\_expert\_num, K\)。
184 183 
185- - H:表示hidden size隐藏层大小。184+ - H:表示hidden size隐藏层大小。
186- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`H`的取值范围如下所示。185+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`H`的取值范围如下所示。
187 - `comm_alg`设置为"fullmesh"时,`H`的取值范围\(0, 7168\],且保证是32的整数倍。186 - `comm_alg`设置为"fullmesh"时,`H`的取值范围\(0, 7168\],且保证是32的整数倍。
188 - `comm_alg`设置为"hierarchy"且驱动版本不低于25.0.RC1.1时,`H`的取值范围\(0, 10 * 1024\],且保证是32的整数倍。187 - `comm_alg`设置为"hierarchy"且驱动版本不低于25.0.RC1.1时,`H`的取值范围\(0, 10 * 1024\],且保证是32的整数倍。
189- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[1024, 8192]。188+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[1024, 8192]。
190 189 
191- - BS:表示待发送的token数量。190+ - BS:表示待发送的token数量。
192- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`BS`的取值范围如下所示。191+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`BS`的取值范围如下所示。
193 - `comm_alg`设置为"fullmesh"时,`BS`的取值范围为0<BS≤256。192 - `comm_alg`设置为"fullmesh"时,`BS`的取值范围为0<BS≤256。
194 - `comm_alg`设置为"hierarchy"且Ascend HDK版本不低于25.0.RC1.1时,`BS`的取值范围为0<BS≤512。193 - `comm_alg`设置为"hierarchy"且Ascend HDK版本不低于25.0.RC1.1时,`BS`的取值范围为0<BS≤512。
195- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`BS`的取值范围如下所示。194+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`BS`的取值范围如下所示。
196 - `comm_alg`设置为""时,`BS`的取值范围为0<BS≤512。195 - `comm_alg`设置为""时,`BS`的取值范围为0<BS≤512。
197 - `comm_alg`设置为"hierarchy"时,`BS`的取值范围为0<BS≤256。196 - `comm_alg`设置为"hierarchy"时,`BS`的取值范围为0<BS≤256。
198 197 
199- - K:表示选取topK个专家,取值范围为0<K≤16,同时满足0 < K ≤ moe\_expert\_num + zero_expert_num + copy_expert_num + const_expert_num。198+ - K:表示选取topK个专家,取值范围为0<K≤16,同时满足0 < K ≤ moe\_expert\_num + zero_expert_num + copy_expert_num + const_expert_num。
200 199 
201- - server\_num:表示服务器的节点数,取值只支持2、4、8。200+ - server\_num:表示服务器的节点数,取值只支持2、4、8。
202- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:仅该场景的shape使用了该变量。201+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:仅该场景的shape使用了该变量。
203 202 
204- - local\_expert\_num:表示本卡专家数量。203+ - local\_expert\_num:表示本卡专家数量。
205- - 对于共享专家卡,local\_expert\_num=1。204+ - 对于共享专家卡,local\_expert\_num=1。
206- - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num),当local\_expert\_num\>1时,不支持TP域通信。205+ - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num),当local\_expert\_num\>1时,不支持TP域通信。
207- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:应满足0 < local\_expert\_num * ep\_world\_size ≤ 2048。206+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:应满足0 < local\_expert\_num * ep\_world\_size ≤ 2048。
208 207 
209-- HCCL通信域缓存区大小:208+- HCCL通信域缓存区大小:
210 209 
211 调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。210 调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。
212- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:211+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
213 该场景支持通过环境变量HCCL\_BUFFSIZE配置。212 该场景支持通过环境变量HCCL\_BUFFSIZE配置。
214 - `comm_alg`配置为"": 仅在此配置下HCCL\_INTRA\_PCIE\_ENABLE和HCCL\_INTRA\_ROCE\_ENABLE生效,依照HCCL\_INTRA\_PCIE\_ENABLE和HCCL\_INTRA\_ROCE\_ENABLE配置选择"fullmesh"或"hierarchy"公式。213 - `comm_alg`配置为"": 仅在此配置下HCCL\_INTRA\_PCIE\_ENABLE和HCCL\_INTRA\_ROCE\_ENABLE生效,依照HCCL\_INTRA\_PCIE\_ENABLE和HCCL\_INTRA\_ROCE\_ENABLE配置选择"fullmesh"或"hierarchy"公式。
215 - `comm_alg`配置为"fullmesh": 设置大小要求\>=2\*\(BS\*ep\_world\_size\*min\(local\_expert\_num, K\)\*H\*sizeof\(uint16\)+2MB\)。214 - `comm_alg`配置为"fullmesh": 设置大小要求\>=2\*\(BS\*ep\_world\_size\*min\(local\_expert\_num, K\)\*H\*sizeof\(uint16\)+2MB\)。
216 - `comm_alg`配置为"hierarchy": 设置大小要求 \>= \(moe\_expert\_num + ep\_world\_size / 4\) \* Align512\(max_bs \* \(H \* sizeof\(dtype_x\) + 4 \* Align8\(K\) \* sizeof\(uint32\)\)\) \* 1B + 8MB,其中Align512\(x\) = \(\(x+512-1\)/512\)\*512,Align8\(x\) = \(\(x+8-1\)/8\)\*8。215 - `comm_alg`配置为"hierarchy": 设置大小要求 \>= \(moe\_expert\_num + ep\_world\_size / 4\) \* Align512\(max_bs \* \(H \* sizeof\(dtype_x\) + 4 \* Align8\(K\) \* sizeof\(uint32\)\)\) \* 1B + 8MB,其中Align512\(x\) = \(\(x+512-1\)/512\)\*512,Align8\(x\) = \(\(x+8-1\)/8\)\*8。
217 216 
218- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:217+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
219 该场景不仅支持通过环境变量HCCL\_BUFFSIZE配置,还支持通过hccl_buffer_size配置(参考《[PyTorch训练模型迁移调优](https://hiascend.com/document/redirect/canncommercial-ptmigr)》中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。218 该场景不仅支持通过环境变量HCCL\_BUFFSIZE配置,还支持通过hccl_buffer_size配置(参考《[PyTorch训练模型迁移调优](https://hiascend.com/document/redirect/canncommercial-ptmigr)》中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。
220 - ep通信域内:设置大小要求 \>= 2且满足\>= 2 \* \(local\_expert\_num \* max\_bs \* ep\_world\_size \* Align512\(Align32\(2 \* h\) + 64\) + \(k + shared\_expert\_num\) \* max\_bs\* Align512\(2 \* h\)\)。219 - ep通信域内:设置大小要求 \>= 2且满足\>= 2 \* \(local\_expert\_num \* max\_bs \* ep\_world\_size \* Align512\(Align32\(2 \* h\) + 64\) + \(k + shared\_expert\_num\) \* max\_bs\* Align512\(2 \* h\)\)。
221 - tp通信域内:设置大小要求 \>= (A \* Align512(Align32(h \* 2) + 44) + A \* Align512(h \* 2)) \* 2。220 - tp通信域内:设置大小要求 \>= (A \* Align512(Align32(h \* 2) + 44) + A \* Align512(h \* 2)) \* 2。
222 - 其中 480Align512(x) = ((x+480-1)/480)\*512,Align512(x) = ((x+512-1)/512)\*512,Align32(x) = ((x+32-1)/32)\*32。221 - 其中 480Align512(x) = ((x+480-1)/480)\*512,Align512(x) = ((x+512-1)/512)\*512,Align32(x) = ((x+32-1)/32)\*32。
223 - `comm_alg`配置为"hierarchy"时,仅支持通过环境变量HCCL\_BUFFSIZE配置。222 - `comm_alg`配置为"hierarchy"时,仅支持通过环境变量HCCL\_BUFFSIZE配置。
224 223 
225-- HCCL_INTRA_PCIE_ENABLE和HCCL_INTRA_ROCE_ENABLE:224+- HCCL_INTRA_PCIE_ENABLE和HCCL_INTRA_ROCE_ENABLE:
226- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:该环境变量不再推荐使用,建议`comm_alg`配置"hierarchy"。225+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:该环境变量不再推荐使用,建议`comm_alg`配置"hierarchy"。
227 226 
228-- 本文公式中的“/”表示整除。227+- 本文公式中的“/”表示整除。
229 228 
230-- 通信域使用约束:229+- 通信域使用约束:
231 230 
232- - 一个模型中的`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`算子仅支持相同EP通信域,且该通信域中不允许有其他算子。231+ - 一个模型中的`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`算子仅支持相同EP通信域,且该通信域中不允许有其他算子。
233 232 
234- - 一个模型中的`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。233+ - 一个模型中的`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。
235 234 
236-- 组网约束:235+- 组网约束:
237- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:多机场景仅支持交换机组网,不支持双机直连组网。236+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:多机场景仅支持交换机组网,不支持双机直连组网。
238 237 
239## 调用示例<a name="zh-cn_topic_0000002168254826_section14459801435"></a>238## 调用示例<a name="zh-cn_topic_0000002168254826_section14459801435"></a>
240 239 
241-- 单算子模式调用240+- 单算子模式调用
242 241 
243 ```python242 ```python
244 import os243 import os
@@ -479,7 +478,7 @@ torch_npu.npu_moe_distribute_combine_v2(expand_x, expert_ids, assist_info_for_co
479 print("run npu success.")478 print("run npu success.")
480 ```479 ```
481 480 
482-- 图模式调用481+- 图模式调用
483 482 
484 ```python483 ```python
485 # 仅支持静态图484 # 仅支持静态图
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_distribute_dispatch.md+93-94
@@ -9,154 +9,154 @@
9 9 
10## 功能说明<a name="zh-cn_topic_0000002203575833_section14441124184110"></a>10## 功能说明<a name="zh-cn_topic_0000002203575833_section14441124184110"></a>
11 11 
12-- API功能:需与[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)配套使用,完成MoE的并行部署下的token dispatch与combine。对token数据先进行quant量化(可选),再进行EP(Expert Parallelism)域的alltoallv通信,再进行TP(Tensor Parallelism)域的allgatherv通信(可选)。12+- API功能:需与[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)配套使用,完成MoE的并行部署下的token dispatch与combine。对token数据先进行quant量化(可选),再进行EP(Expert Parallelism)域的alltoallv通信,再进行TP(Tensor Parallelism)域的allgatherv通信(可选)。
13-- 计算公式:$x$表示输入`x`,$scales$表示输入`scales`,$quant\_mode$表示输入`quant_mode`。13+- 计算公式:$x$表示输入`x`,$scales$表示输入`scales`,$quant\_mode$表示输入`quant_mode`。
14- - 若`quant_mode`不为`2`,即非动态量化场景:14+ - 若`quant_mode`不为`2`,即非动态量化场景:
15 15 
16 ![](../../figures/zh-cn_formulaimage_0000002244554688.png)16 ![](../../figures/zh-cn_formulaimage_0000002244554688.png)
17 17 
18- - 若`quant_mode`为`2`,即动态量化场景:18+ - 若`quant_mode`为`2`,即动态量化场景:
19 19 
20 ![](../../figures/zh-cn_formulaimage_0000002244394892.png)20 ![](../../figures/zh-cn_formulaimage_0000002244394892.png)
21 21 
22## 函数原型<a name="zh-cn_topic_0000002203575833_section45077510411"></a>22## 函数原型<a name="zh-cn_topic_0000002203575833_section45077510411"></a>
23 23 
24-```24+```python
25torch_npu.npu_moe_distribute_dispatch(x, expert_ids, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, scales=None, x_active_mask=None, expert_scales=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, quant_mode=0, global_bs=0, expert_token_nums_type=1) -> (Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor)25torch_npu.npu_moe_distribute_dispatch(x, expert_ids, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, scales=None, x_active_mask=None, expert_scales=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, quant_mode=0, global_bs=0, expert_token_nums_type=1) -> (Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor)
26```26```
27 27 
28## 参数说明<a name="zh-cn_topic_0000002203575833_section112637109429"></a>28## 参数说明<a name="zh-cn_topic_0000002203575833_section112637109429"></a>
29 29 
30-- **x**(`Tensor`):表示计算使用的token数据,需根据`expert_ids`来发送给其他卡。要求2维张量,shape为\(BS, H\),表示有BS(batch size)个token,数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。30+- **x**(`Tensor`):表示计算使用的token数据,需根据`expert_ids`来发送给其他卡。要求2维张量,shape为\(BS, H\),表示有BS(batch size)个token,数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。
31-- **expert\_ids**(`Tensor`):表示每个token的topK个专家索引,决定每个token要发给哪些专家。要求2维张量,shape为\(BS, K\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)的`expert_ids`输入,张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。31+- **expert\_ids**(`Tensor`):表示每个token的topK个专家索引,决定每个token要发给哪些专家。要求2维张量,shape为\(BS, K\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)的`expert_ids`输入,张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。
32-- **group\_ep**(`str`):EP通信域名称,专家并行的通信域。字符串长度范围为\[1,128\)。<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>时不能和`group_tp`相同。32+- **group\_ep**(`str`):EP通信域名称,专家并行的通信域。字符串长度范围为\[1,128\)。<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>时不能和`group_tp`相同。
33-- **ep\_world\_size**(`int`):EP通信域size。33+- **ep\_world\_size**(`int`):EP通信域size。
34- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值支持16、32、64。34+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值支持16、32、64。
35- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持8、16、32、64、128、144、256、288。35+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持8、16、32、64、128、144、256、288。
36 36 
37-- **ep\_rank\_id**(`int`):EP通信域本卡ID,取值范围\[0, ep\_world\_size\),同一个EP通信域中各卡的`ep_rank_id`不重复。37+- **ep\_rank\_id**(`int`):EP通信域本卡ID,取值范围\[0, ep\_world\_size\),同一个EP通信域中各卡的`ep_rank_id`不重复。
38-- **moe\_expert\_num**(`int`):MoE专家数量,取值范围\[1, 512\],并且满足以下条件:moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0。38+- **moe\_expert\_num**(`int`):MoE专家数量,取值范围\[1, 512\],并且满足以下条件:moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0。
39- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:还需满足moe\_expert\_num\/\(ep\_world\_size - shared\_expert\_rank\_num\) <= 24。39+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:还需满足moe\_expert\_num\/\(ep\_world\_size - shared\_expert\_rank\_num\) <= 24。
40- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。40- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
41-- **scales**(`Tensor`):可选参数,表示每个专家的权重,非量化场景不传入,动态量化场景可传可不传。若传值要求为2维张量,如果有共享专家,shape为\(shared\_expert\_num+moe\_expert\_num, H\),如果没有共享专家,shape为\(moe\_expert\_num, H\),数据类型支持`float`,数据格式为$ND$,不支持非连续的Tensor。41+- **scales**(`Tensor`):可选参数,表示每个专家的权重,非量化场景不传入,动态量化场景可传可不传。若传值要求为2维张量,如果有共享专家,shape为\(shared\_expert\_num+moe\_expert\_num, H\),如果没有共享专家,shape为\(moe\_expert\_num, H\),数据类型支持`float`,数据格式为$ND$,不支持非连续的Tensor。
42-- **x\_active\_mask**(`Tensor`):预留参数,暂未使用,使用默认值即可。42+- **x\_active\_mask**(`Tensor`):预留参数,暂未使用,使用默认值即可。
43 43 
44-- **expert\_scales**(`Tensor`):可选参数,表示每个token的topK个专家权重。44+- **expert\_scales**(`Tensor`):可选参数,表示每个token的topK个专家权重。
45- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求2维张量,shape为\(BS, K\),数据类型支持`float`,数据格式为$ND$,支持非连续的Tensor。45+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求2维张量,shape为\(BS, K\),数据类型支持`float`,数据格式为$ND$,支持非连续的Tensor。
46- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:暂不支持该参数,使用默认值即可。46+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:暂不支持该参数,使用默认值即可。
47 47 
48-- **group\_tp**(`str`):可选参数,TP通信域名称,数据并行的通信域。若有TP域通信需要传参,若无TP域通信,使用默认值""即可。48+- **group\_tp**(`str`):可选参数,TP通信域名称,数据并行的通信域。若有TP域通信需要传参,若无TP域通信,使用默认值""即可。
49- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:eager模式使用默认值即可,图模式传入与`group_ep`相同。49+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:eager模式使用默认值即可,图模式传入与`group_ep`相同。
50- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。50+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。
51 51 
52-- **tp\_world\_size**(`int`):可选参数,TP通信域size。有TP域通信才需要传参。52+- **tp\_world\_size**(`int`):可选参数,TP通信域size。有TP域通信才需要传参。
53- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。53+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。
54- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。54+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。
55 55 
56-- **tp\_rank\_id**(`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。56+- **tp\_rank\_id**(`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。
57- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值即可。57+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值即可。
58- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],默认为0,同一个TP通信域中各卡的`tp_rank_id`不重复。无TP域通信时,传0即可。58+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],默认为0,同一个TP通信域中各卡的`tp_rank_id`不重复。无TP域通信时,传0即可。
59 59 
60-- **expert\_shard\_type**(`int`):表示共享专家卡排布类型。60+- **expert\_shard\_type**(`int`):表示共享专家卡排布类型。
61- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。61+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。
62- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。62+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。
63-- **shared\_expert\_num**(`int`):表示共享专家数量,一个共享专家可以复制部署到多个卡上。63+- **shared\_expert\_num**(`int`):表示共享专家数量,一个共享专家可以复制部署到多个卡上。
64- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。64+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。
65- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅支持1,默认值为1。65+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅支持1,默认值为1。
66 66 
67-- **shared\_expert\_rank\_num**(`int`):可选参数,表示共享专家卡数量。67+- **shared\_expert\_rank\_num**(`int`):可选参数,表示共享专家卡数量。
68- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家,传0即可。68+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家,传0即可。
69- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, ep\_world\_size\)。取0表示无共享专家,不取0时需满足ep\_world\_size%shared\_expert\_rank\_num=0。69+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, ep\_world\_size\)。取0表示无共享专家,不取0时需满足ep\_world\_size%shared\_expert\_rank\_num=0。
70 70 
71-- **quant\_mode**(`int`):可选参数,表示量化模式。支持取值:0表示非量化(默认),2表示动态量化。当`quant_mode`为2,`dynamic_scales`不为None;当`quant_mode`为0,`dynamic_scales`为None。71+- **quant\_mode**(`int`):可选参数,表示量化模式。支持取值:0表示非量化(默认),2表示动态量化。当`quant_mode`为2,`dynamic_scales`不为None;当`quant_mode`为0,`dynamic_scales`为None。
72-- **global\_bs**(`int`):可选参数,表示EP域全局的BS大小。72+- **global\_bs**(`int`):可选参数,表示EP域全局的BS大小。
73- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size或者256\*ep\_world\_size,其中max\_bs表示单rank BS最大值,建议按max\_bs\*ep\_world\_size传入,固定按256\*ep\_world\_size传入,在后续版本BS大于256的场景下会无法支持;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。73+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size或者256\*ep\_world\_size,其中max\_bs表示单rank BS最大值,建议按max\_bs\*ep\_world\_size传入,固定按256\*ep\_world\_size传入,在后续版本BS大于256的场景下会无法支持;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。
74- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。74+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。
75 75 
76-- **expert\_token\_nums\_type**(`int`):可选参数,表示输出`expert_token_nums`的值类型,取值范围\[0, 1\],0表示每个专家收到token数量的前缀和,1表示每个专家收到的token数量(默认)。76+- **expert\_token\_nums\_type**(`int`):可选参数,表示输出`expert_token_nums`的值类型,取值范围\[0, 1\],0表示每个专家收到token数量的前缀和,1表示每个专家收到的token数量(默认)。
77 77 
78## 返回值说明<a name="zh-cn_topic_0000002203575833_section22231435517"></a>78## 返回值说明<a name="zh-cn_topic_0000002203575833_section22231435517"></a>
79 79 
80-- **expand\_x**(`Tensor`):表示本卡收到的token数据,要求2维张量,shape为\(max\(tp\_world\_size, 1\) \*A, H\),A表示在EP通信域可能收到的最大token数,数据类型支持`bfloat16`、`float16`、`int8`。量化时类型为`int8`,非量化时与`x`数据类型保持一致。数据格式为$ND$,支持非连续的Tensor。80+- **expand\_x**(`Tensor`):表示本卡收到的token数据,要求2维张量,shape为\(max\(tp\_world\_size, 1\) \*A, H\),A表示在EP通信域可能收到的最大token数,数据类型支持`bfloat16`、`float16`、`int8`。量化时类型为`int8`,非量化时与`x`数据类型保持一致。数据格式为$ND$,支持非连续的Tensor。
81-- **dynamic\_scales**(`Tensor`):表示计算得到的动态量化参数。当`quant_mode`非0时才有该输出,要求1维张量,shape为\(A,\),数据类型支持`float`,数据格式支持$ND$,支持非连续的Tensor。81+- **dynamic\_scales**(`Tensor`):表示计算得到的动态量化参数。当`quant_mode`非0时才有该输出,要求1维张量,shape为\(A,\),数据类型支持`float`,数据格式支持$ND$,支持非连续的Tensor。
82-- **expand\_idx**(`Tensor`):表示给同一专家发送的token个数,要求1维张量,shape为\(BS \* K, \)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)的`expand_idx`输入。82+- **expand\_idx**(`Tensor`):表示给同一专家发送的token个数,要求1维张量,shape为\(BS \* K, \)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)的`expand_idx`输入。
83 83 
84-- **expert\_token\_nums**(`Tensor`):本卡每个专家实际收到的token数量,要求1维张量,shape为\(local\_expert\_num,\),数据类型`int64`,数据格式支持$ND$,支持非连续的Tensor。84+- **expert\_token\_nums**(`Tensor`):本卡每个专家实际收到的token数量,要求1维张量,shape为\(local\_expert\_num,\),数据类型`int64`,数据格式支持$ND$,支持非连续的Tensor。
85-- **ep\_recv\_counts**(`Tensor`):表示EP通信域各卡收到的token数(token数以前缀和的形式表示),要求1维张量,数据类型`int32`,数据格式支持$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)的`ep_send_counts`输入。85+- **ep\_recv\_counts**(`Tensor`):表示EP通信域各卡收到的token数(token数以前缀和的形式表示),要求1维张量,数据类型`int32`,数据格式支持$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)的`ep_send_counts`输入。
86- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(moe\_expert\_num+2\*global\_bs\*K\*server\_num, \),前`moe_expert_num`个数表示在EP通信域内,该卡上每个专家收到来自其他各卡的token数(以前缀和的形式表示),2\*global\_bs\*K\*server\_num用于存储机间和机内通信前,combine可提前做reduce操作的token个数和通信区偏移量,`global_bs`传入0时此处按照bs\*ep\_world\_size计算。86+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(moe\_expert\_num+2\*global\_bs\*K\*server\_num, \),前`moe_expert_num`个数表示在EP通信域内,该卡上每个专家收到来自其他各卡的token数(以前缀和的形式表示),2\*global\_bs\*K\*server\_num用于存储机间和机内通信前,combine可提前做reduce操作的token个数和通信区偏移量,`global_bs`传入0时此处按照bs\*ep\_world\_size计算。
87- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(ep\_world\_size\*max\(tp\_world\_size, 1\)\*local\_expert\_num, \)。87+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(ep\_world\_size\*max\(tp\_world\_size, 1\)\*local\_expert\_num, \)。
88 88 
89-- **tp\_recv\_counts**(`Tensor`):表示TP通信域各卡收到的token数量。对应[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)的`tp_send_counts`输入。89+- **tp\_recv\_counts**(`Tensor`):表示TP通信域各卡收到的token数量。对应[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)的`tp_send_counts`输入。
90- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP通信域,暂无该输出,90+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP通信域,暂无该输出,
91- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求一个1维张量,shape为\(tp\_world\_size, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。91+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求一个1维张量,shape为\(tp\_world\_size, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。
92 92 
93-- **expand\_scales**(`Tensor`):表示`expert_scales`与`x`一起进行alltoallv之后的输出。93+- **expand\_scales**(`Tensor`):表示`expert_scales`与`x`一起进行alltoallv之后的输出。
94- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求一个1维张量,shape为\(A, \),数据类型支持`float`,数据格式要求为$ND$,支持非连续的Tensor。94+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求一个1维张量,shape为\(A, \),数据类型支持`float`,数据格式要求为$ND$,支持非连续的Tensor。
95- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:暂不支持该输出,返回None。95+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:暂不支持该输出,返回None。
96 96 
97## 约束说明<a name="zh-cn_topic_0000002203575833_section12345537164214"></a>97## 约束说明<a name="zh-cn_topic_0000002203575833_section12345537164214"></a>
98 98 
99-- 该接口支持推理场景下使用。99+- 该接口支持推理场景下使用。
100-- 该接口支持静态图模式,`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`必须配套使用。100+- 该接口支持静态图模式,`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`必须配套使用。
101-- 在不同产品型号、不同通信算法或不同版本中,`npu_moe_distribute_dispatch`的Tensor输出`expand_idx`、`ep_recv_counts`、`tp_recv_counts`、`expand_scales`中的元素值可能不同,使用时直接将上述Tensor传给`npu_moe_distribute_combine`对应参数即可,模型其他业务逻辑不应对其存在依赖。101+- 在不同产品型号、不同通信算法或不同版本中,`npu_moe_distribute_dispatch`的Tensor输出`expand_idx`、`ep_recv_counts`、`tp_recv_counts`、`expand_scales`中的元素值可能不同,使用时直接将上述Tensor传给`npu_moe_distribute_combine`对应参数即可,模型其他业务逻辑不应对其存在依赖。
102-- 调用接口过程中使用的`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`shared_expert_num`、`shared_expert_rank_num`、`global_bs`参数取值所有卡需保持一致,`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`global_bs`网络中不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)对应参数也保持一致。102+- 调用接口过程中使用的`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`shared_expert_num`、`shared_expert_rank_num`、`global_bs`参数取值所有卡需保持一致,`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`global_bs`网络中不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_combine](torch_npu-npu_moe_distribute_combine.md)对应参数也保持一致。
103-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。103+- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。
104-- 参数里Shape使用的变量如下:104+- 参数里Shape使用的变量如下:
105- - A:表示本卡接收的最大token数量,取值范围如下105+ - A:表示本卡接收的最大token数量,取值范围如下
106- - 对于共享专家,当global\_bs为0时,要满足A=BS\*shared\_expert\_num/shared\_expert\_rank\_num;当global\_bs非0时,要满足A=global\_bs\*shared\_expert\_num/shared\_expert\_rank\_num。106+ - 对于共享专家,当global\_bs为0时,要满足A=BS\*shared\_expert\_num/shared\_expert\_rank\_num;当global\_bs非0时,要满足A=global\_bs\*shared\_expert\_num/shared\_expert\_rank\_num。
107- - 对于MoE专家,当global\_bs为0时,要满足A\>=BS\*ep\_world\_size\*min\(local\_expert\_num, K\);当global\_bs非0时,要满足A\>=global\_bs\* min\(local\_expert\_num, K\)。107+ - 对于MoE专家,当global\_bs为0时,要满足A\>=BS\*ep\_world\_size\*min\(local\_expert\_num, K\);当global\_bs非0时,要满足A\>=global\_bs\* min\(local\_expert\_num, K\)。
108 108 
109- - H:表示hidden size隐藏层大小。109+ - H:表示hidden size隐藏层大小。
110- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值范围\(0,7168\],且保证是32的整数倍。110+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值范围\(0,7168\],且保证是32的整数倍。
111- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅支持 7168。111+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:仅支持 7168。
112 112 
113- - BS:表示待发送的token数量。113+ - BS:表示待发送的token数量。
114- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值范围为0<BS≤256。114+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值范围为0<BS≤256。
115- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为0<BS≤512。115+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为0<BS≤512。
116 116 
117- - K:表示选取topK个专家,需满足0<K≤moe\_expert\_num。117+ - K:表示选取topK个专家,需满足0<K≤moe\_expert\_num。
118- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:保证取值范围为0<K≤16。118+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:保证取值范围为0<K≤16。
119- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:保证取值范围为0<K≤8。119+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:保证取值范围为0<K≤8。
120 120 
121- - server\_num:表示服务器的节点数,取值只支持2、4、8。121+ - server\_num:表示服务器的节点数,取值只支持2、4、8。
122- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:仅该场景的shape使用了该变量。122+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:仅该场景的shape使用了该变量。
123 123 
124- - local\_expert\_num:表示本卡专家数量。124+ - local\_expert\_num:表示本卡专家数量。
125- - 对于共享专家卡,local\_expert\_num=1125+ - 对于共享专家卡,local\_expert\_num=1
126- - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num),当local\_expert\_num\>1时,不支持TP域通信。126+ - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num),当local\_expert\_num\>1时,不支持TP域通信。
127 127 
128-- HCCL通信域缓存区大小:128+- HCCL通信域缓存区大小:
129 129 
130 调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。130 调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。
131- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:131+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
132 该场景支持通过环境变量HCCL\_BUFFSIZE配置。132 该场景支持通过环境变量HCCL\_BUFFSIZE配置。
133 - 设置大小要求\>=2\*\(BS\*ep\_world\_size\*min\(local\_expert\_num, K\)\*H\*sizeof\(uint16\)+2MB\)。133 - 设置大小要求\>=2\*\(BS\*ep\_world\_size\*min\(local\_expert\_num, K\)\*H\*sizeof\(uint16\)+2MB\)。
134- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:134+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
135 该场景不仅支持通过环境变量HCCL\_BUFFSIZE配置,还支持通过hccl_buffer_size配置(参考《[PyTorch训练模型迁移调优](https://hiascend.com/document/redirect/canncommercial-ptmigr)》中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。135 该场景不仅支持通过环境变量HCCL\_BUFFSIZE配置,还支持通过hccl_buffer_size配置(参考《[PyTorch训练模型迁移调优](https://hiascend.com/document/redirect/canncommercial-ptmigr)》中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。
136 - ep通信域内:设置大小要求\>=2且满足1024\^2\*\(HCCL\_BUFFSIZE\-2\)\/2\>=BS\*2\*\(H\+128\)\*\(ep\_world\_size\*local\_expert\_num\+K\+1\),local\_expert\_num需使用MoE专家卡的本卡专家数。136 - ep通信域内:设置大小要求\>=2且满足1024\^2\*\(HCCL\_BUFFSIZE\-2\)\/2\>=BS\*2\*\(H\+128\)\*\(ep\_world\_size\*local\_expert\_num\+K\+1\),local\_expert\_num需使用MoE专家卡的本卡专家数。
137- - tp通信域内:设置大小要求\>=A * (H * 2 + 128) * 2。137+ - tp通信域内:设置大小要求\>=A \* (H \* 2 + 128) \* 2。
138-- <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:配置环境变量HCCL\_INTRA\_PCIE\_ENABLE=1和HCCL\_INTRA\_ROCE\_ENABLE=0可以减少跨机通信数据量,提升算子性能。此时要求HCCL\_BUFFSIZE\>=moe\_expert\_num\*BS\*\(H\*sizeof\(dtype_x\)+4\*\(\(K+7\)/8\*8\)\*sizeof\(uint32\)\)+4MB+100MB。并且,对于入参moe\_expert\_num,只要求moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0,不要求moe\_expert\_num\/\(ep\_world\_size - shared\_expert\_rank\_num\) <= 24。138+- <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:配置环境变量HCCL\_INTRA\_PCIE\_ENABLE=1和HCCL\_INTRA\_ROCE\_ENABLE=0可以减少跨机通信数据量,提升算子性能。此时要求HCCL\_BUFFSIZE\>=moe\_expert\_num\*BS\*\(H\*sizeof\(dtype_x\)+4\*\(\(K+7\)/8\*8\)\*sizeof\(uint32\)\)+4MB+100MB。并且,对于入参moe\_expert\_num,只要求moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0,不要求moe\_expert\_num\/\(ep\_world\_size - shared\_expert\_rank\_num\) <= 24。
139 139 
140-- 本文公式中的“/”表示整除。140+- 本文公式中的“/”表示整除。
141 141 
142-- 通信域使用约束:142+- 通信域使用约束:
143 143 
144- - 一个模型中的`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`算子仅支持相同EP通信域,且该通信域中不允许有其他算子。144+ - 一个模型中的`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`算子仅支持相同EP通信域,且该通信域中不允许有其他算子。
145 145 
146- - 一个模型中的`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。146+ - 一个模型中的`npu_moe_distribute_dispatch`和`npu_moe_distribute_combine`算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。
147 147 
148- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:一个通信域内的节点需在一个超节点内,不支持跨超节点。148+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:一个通信域内的节点需在一个超节点内,不支持跨超节点。
149 149 
150-- 组网约束:150+- 组网约束:
151- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:多机场景仅支持交换机组网,不支持双机直连组网。151+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:多机场景仅支持交换机组网,不支持双机直连组网。
152 152 
153-- 版本配套约束:153+- 版本配套约束:
154 154 
155 静态图模式下,从Ascend Extension for PyTorch 8.0.0版本开始,Ascend Extension for PyTorch框架会对静态图中最后一个节点输出结果做Meta推导与inferShape推导的结果强校验。当图中只有一个Dispatch算子,若CANN版本落后于Ascend Extension for PyTorch版本,会出现Shape不匹配报错,建议用户升级CANN版本,详细的版本配套关系参见《[Ascend Extension for PyTorch 版本说明](https://gitcode.com/Ascend/pytorch/blob/v2.7.1-7.3.0/docs/zh/release_notes/release_notes.md)》中“相关产品版本配套说明””。155 静态图模式下,从Ascend Extension for PyTorch 8.0.0版本开始,Ascend Extension for PyTorch框架会对静态图中最后一个节点输出结果做Meta推导与inferShape推导的结果强校验。当图中只有一个Dispatch算子,若CANN版本落后于Ascend Extension for PyTorch版本,会出现Shape不匹配报错,建议用户升级CANN版本,详细的版本配套关系参见《[Ascend Extension for PyTorch 版本说明](https://gitcode.com/Ascend/pytorch/blob/v2.7.1-7.3.0/docs/zh/release_notes/release_notes.md)》中“相关产品版本配套说明””。
156 156 
157## 调用示例<a name="zh-cn_topic_0000002203575833_section14459801435"></a>157## 调用示例<a name="zh-cn_topic_0000002203575833_section14459801435"></a>
158 158 
159-- 单算子模式调用159+- 单算子模式调用
160 160 
161 ```python161 ```python
162 import os162 import os
@@ -319,7 +319,7 @@ torch_npu.npu_moe_distribute_dispatch(x, expert_ids, group_ep, ep_world_size, ep
319 319 
320 ```320 ```
321 321 
322-- 图模式调用322+- 图模式调用
323 323 
324 ```python324 ```python
325 # 仅支持静态图325 # 仅支持静态图
@@ -501,4 +501,3 @@ torch_npu.npu_moe_distribute_dispatch(x, expert_ids, group_ep, ep_world_size, ep
501 p.join()501 p.join()
502 print("run npu success.")502 print("run npu success.")
503 ```503 ```
504- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_distribute_dispatch_v2.md+107-109
@@ -9,16 +9,16 @@
9 9 
10## 功能说明<a name="zh-cn_topic_0000002203575833_section14441124184110"></a>10## 功能说明<a name="zh-cn_topic_0000002203575833_section14441124184110"></a>
11 11 
12-- API功能:12+- API功能:
13 需与[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)或[torch\_npu.npu\_moe\_distribute\_combine\_add\_rms\_norm](torch_npu-npu_moe_distribute_combine_add_rms_norm.md)配套使用,完成MoE的并行部署下的token dispatch\_v2与combine\_v2。13 需与[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)或[torch\_npu.npu\_moe\_distribute\_combine\_add\_rms\_norm](torch_npu-npu_moe_distribute_combine_add_rms_norm.md)配套使用,完成MoE的并行部署下的token dispatch\_v2与combine\_v2。
14 - 支持动态量化场景,对token数据先进行量化(可选),再进行EP(Expert Parallelism)域的alltoallv通信,再进行TP(Tensor Parallelism)域的allgatherv通信(可选);14 - 支持动态量化场景,对token数据先进行量化(可选),再进行EP(Expert Parallelism)域的alltoallv通信,再进行TP(Tensor Parallelism)域的allgatherv通信(可选);
15 - 支持特殊专家场景。15 - 支持特殊专家场景。
16 16 
17-- 相较于npu_moe_distribute_dispatch接口,该接口变更如下:17+- 相较于npu_moe_distribute_dispatch接口,该接口变更如下:
18- - npu_moe_distribute_dispatch中shape为(Bs * K,)的返回值`expand_idx`替换为shape为(A * 128,)的`assist_info_for_combine`,以包含更详细的token信息辅助torch_npu.npu_moe_distribute_combine_v2高效地进行全卡同步;18+ - npu_moe_distribute_dispatch中shape为(Bs \* K,)的返回值`expand_idx`替换为shape为(A \* 128,)的`assist_info_for_combine`,以包含更详细的token信息辅助torch_npu.npu_moe_distribute_combine_v2高效地进行全卡同步;
19 - 新增输入参数`comm_alg`,可用于代替HCCL_INTRA_PCIE_ENABLE和HCCL_INTRA_ROCE_ENABLE环境变量。19 - 新增输入参数`comm_alg`,可用于代替HCCL_INTRA_PCIE_ENABLE和HCCL_INTRA_ROCE_ENABLE环境变量。
20 20 
21-- 计算公式:21+- 计算公式:
22 - 动态量化场景:22 - 动态量化场景:
23 23 
24 若`quant_mode`不为`2`,即非动态量化场景:24 若`quant_mode`不为`2`,即非动态量化场景:
@@ -69,7 +69,6 @@
69 \end{cases}69 \end{cases}
70 $$70 $$
71 71 
72- 
73 - 特殊专家场景:72 - 特殊专家场景:
74 73 
75 零专家场景,即`zero_expert_num`不为0:74 零专家场景,即`zero_expert_num`不为0:
@@ -88,164 +87,163 @@
88 87 
89## 函数原型<a name="zh-cn_topic_0000002203575833_section45077510411"></a>88## 函数原型<a name="zh-cn_topic_0000002203575833_section45077510411"></a>
90 89 
91-```90+```python
92torch_npu.npu_moe_distribute_dispatch_v2(x, expert_ids, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, scales=None, x_active_mask=None, expert_scales=None, elastic_info=None, performance_info=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, quant_mode=0, global_bs=0, expert_token_nums_type=1, comm_alg="", zero_expert_num=0, copy_expert_num=0, const_expert_num=0) -> (Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor)91torch_npu.npu_moe_distribute_dispatch_v2(x, expert_ids, group_ep, ep_world_size, ep_rank_id, moe_expert_num, *, scales=None, x_active_mask=None, expert_scales=None, elastic_info=None, performance_info=None, group_tp="", tp_world_size=0, tp_rank_id=0, expert_shard_type=0, shared_expert_num=1, shared_expert_rank_num=0, quant_mode=0, global_bs=0, expert_token_nums_type=1, comm_alg="", zero_expert_num=0, copy_expert_num=0, const_expert_num=0) -> (Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor)
93```92```
94 93 
95## 参数说明<a name="zh-cn_topic_0000002203575833_section112637109429"></a>94## 参数说明<a name="zh-cn_topic_0000002203575833_section112637109429"></a>
96 95 
97-- **x** (`Tensor`):必选参数,表示计算使用的token数据,需根据`expert_ids`来发送给其他卡。要求为2维张量,shape为\(BS, H\),表示有BS个token,数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。96+- **x** (`Tensor`):必选参数,表示计算使用的token数据,需根据`expert_ids`来发送给其他卡。要求为2维张量,shape为\(BS, H\),表示有BS个token,数据类型支持`bfloat16`、`float16`,数据格式为$ND$,支持非连续的Tensor。
98-- **expert\_ids** (`Tensor`):必选参数,表示每个token的topK个专家索引,决定每个token要发给哪些专家。要求为2维张量,shape为\(BS, K\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)的`expert_ids`输入,张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。97+- **expert\_ids** (`Tensor`):必选参数,表示每个token的topK个专家索引,决定每个token要发给哪些专家。要求为2维张量,shape为\(BS, K\),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)的`expert_ids`输入,张量里value取值范围为\[0, moe\_expert\_num\),且同一行中的K个value不能重复。
99-- **group\_ep** (`str`):必选参数,EP通信域名称,专家并行的通信域。字符串长度范围为\[1,128\)。<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>时,`group_ep`不能与`group_tp`相同。98+- **group\_ep** (`str`):必选参数,EP通信域名称,专家并行的通信域。字符串长度范围为\[1,128\)。<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>时,`group_ep`不能与`group_tp`相同。
100-- **ep\_world\_size**(`int`):必选参数,EP通信域size。99+- **ep\_world\_size**(`int`):必选参数,EP通信域size。
101- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`ep_world_size`的取值范围如下所示。100+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`ep_world_size`的取值范围如下所示。
102 - `comm_alg`设置为"fullmesh"时,`ep_world_size`取值范围为2、3、4、5、6、7、8、16、32、64、128、256。101 - `comm_alg`设置为"fullmesh"时,`ep_world_size`取值范围为2、3、4、5、6、7、8、16、32、64、128、256。
103 - `comm_alg`设置为"hierarchy"时,`ep_world_size`取值范围为16、32、64。102 - `comm_alg`设置为"hierarchy"时,`ep_world_size`取值范围为16、32、64。
104- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持\[2, 768\]。`comm_alg`设置为"hierarchy"时,取值范围为[16, 256],且为16的整数倍。103+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值支持\[2, 768\]。`comm_alg`设置为"hierarchy"时,取值范围为[16, 256],且为16的整数倍。
105 104 
106- 105+- **ep\_rank\_id** (`int`):必选参数,EP通信域本卡ID,取值范围\[0, ep\_world\_size\),同一个EP通信域中各卡的`ep_rank_id`不重复。
107-- **ep\_rank\_id** (`int`):必选参数,EP通信域本卡ID取值范围\[0, ep\_world\_size\),同一个EP通信域中各卡的`ep_rank_id`不重复106+- **moe\_expert\_num** (`int`):必选参数,MoE专家数量并且满足以下条件:moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0
108-- **moe\_expert\_num** (`int`)必选参数,MoE专家数量,并且满足以下条件:moe\_expert\_num\%\(ep\_world\_size - shared\_expert\_rank\_num\)\=0107+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>取值范围为\[1, 1024\]
109- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:取值范围为\[1, 1024\]。108+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为\[1, 1024\]。`comm_alg`设置为"hierarchy"时,取值范围为[1, 512]。
110- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为\[1, 1024\]。`comm_alg`设置为"hierarchy"时,取值范围为[1, 512]。
111- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。109- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
112-- **scales** (`Tensor`):可选参数,表示每个专家的权重,非量化场景不传,动态量化场景可传可不传。若传值要求为2维张量,如果有共享专家,shape为\(shared\_expert\_num+moe\_expert\_num, H\),如果没有共享专家,shape为\(moe\_expert\_num, H\),数据类型支持`float32`,数据格式为$ND$,不支持非连续的Tensor。110+- **scales** (`Tensor`):可选参数,表示每个专家的权重,非量化场景不传,动态量化场景可传可不传。若传值要求为2维张量,如果有共享专家,shape为\(shared\_expert\_num+moe\_expert\_num, H\),如果没有共享专家,shape为\(moe\_expert\_num, H\),数据类型支持`float32`,数据格式为$ND$,不支持非连续的Tensor。
113-- **x\_active\_mask** (`Tensor`):可选参数,表示token是否参与通信。111+- **x\_active\_mask** (`Tensor`):可选参数,表示token是否参与通信。
114- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:112+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
115 - `comm_alg`设置为"fullmesh"时,要求是一个1维或者2维张量。当输入为1维时,shape为\(BS, \); 当输入为2维时,shape为\(BS, K\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。当输入为1维时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2D时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。支持2维张量属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。113 - `comm_alg`设置为"fullmesh"时,要求是一个1维或者2维张量。当输入为1维时,shape为\(BS, \); 当输入为2维时,shape为\(BS, K\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。当输入为1维时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2D时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。支持2维张量属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。
116 - `comm_alg`设置为"hierarchy"时,当前版本不支持,使用默认值None即可。114 - `comm_alg`设置为"hierarchy"时,当前版本不支持,使用默认值None即可。
117- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求是一个1维或者2维张量。当输入为1维时,shape为\(BS, \); 当输入为2维时,shape为\(BS, K\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。当输入为1维时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2D时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。115+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求是一个1维或者2维张量。当输入为1维时,shape为\(BS, \); 当输入为2维时,shape为\(BS, K\)。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。当输入为1维时,参数为true表示对应的token参与通信,true必须排到false之前,例:{true, false, true} 为非法输入;当输入为2D时,参数为true表示当前token对应的`expert_ids`参与通信,若当前token对应的K个`bool`值全为false,表示当前token不会参与通信。默认所有token都会参与通信。当每张卡的BS数量不一致时,所有token必须全部有效。
118 116 
119-- **expert\_scales** (`Tensor`):可选参数,表示每个token的topK个专家权重。117+- **expert\_scales** (`Tensor`):可选参数,表示每个token的topK个专家权重。
120- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求为2维张量,shape为\(BS, K\),数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。118+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求为2维张量,shape为\(BS, K\),数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。
121- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`comm_alg`设置为"hierarchy"时,要求为2维张量,shape为\(BS, K\),数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。`comm_alg`设置为"","fullmesh_v1","fullmesh_v2"时,暂不支持该参数,使用默认值即可。119+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`comm_alg`设置为"hierarchy"时,要求为2维张量,shape为\(BS, K\),数据类型支持`float32`,数据格式为$ND$,支持非连续的Tensor。`comm_alg`设置为"","fullmesh_v1","fullmesh_v2"时,暂不支持该参数,使用默认值即可。
122 120 
123-- **elastic\_info** (`Tensor`):预留参数,当前版本不支持,传默认值None即可。121+- **elastic\_info** (`Tensor`):预留参数,当前版本不支持,传默认值None即可。
124 122 
125-- **performance\_info** (`Tensor`):可选参数,表示本卡等待各卡数据的通信时间,单位为us(微秒)。单次算子调用各卡通信耗时会累加到该Tensor上,算子内部不进行自动清零,因此每次启用此Tensor开始记录耗时前需对Tensor清零。当传入None时表示不使能记录通信耗时功能;当传入有效数据时,要求是一个1D的Tensor,shape为(ep_world_size,),数据类型支持int64,数据格式要求为ND,支持非连续的Tensor。123+- **performance\_info** (`Tensor`):可选参数,表示本卡等待各卡数据的通信时间,单位为us(微秒)。单次算子调用各卡通信耗时会累加到该Tensor上,算子内部不进行自动清零,因此每次启用此Tensor开始记录耗时前需对Tensor清零。当传入None时表示不使能记录通信耗时功能;当传入有效数据时,要求是一个1D的Tensor,shape为(ep_world_size,),数据类型支持int64,数据格式要求为ND,支持非连续的Tensor。
126 124 
127-- **group\_tp** (`string`):可选参数,TP通信域名称,数据并行的通信域。若有TP域通信需要传参,若无TP域通信,使用默认值""即可。125+- **group\_tp** (`string`):可选参数,TP通信域名称,数据并行的通信域。若有TP域通信需要传参,若无TP域通信,使用默认值""即可。
128- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:eager模式使用默认值即可,图模式传入与`group_ep`相同。126+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:eager模式使用默认值即可,图模式传入与`group_ep`相同。
129- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。127+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:字符串长度范围为\[0, 128\),不能和`group_ep`相同,仅在无TP域时支持传空。
130 128 
131-- **tp\_world\_size** (`int`):可选参数,TP通信域size。有TP域通信才需要传参。129+- **tp\_world\_size** (`int`):可选参数,TP通信域size。有TP域通信才需要传参。
132- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。130+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值0即可。
133- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。131+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 2\],0和1表示无TP域通信,2表示有TP域通信。
134 132 
135-- **tp\_rank\_id** (`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。133+- **tp\_rank\_id** (`int`):可选参数,TP通信域本卡ID。有TP域通信才需要传参。
136- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值即可。134+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP域通信,使用默认值即可。
137- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],默认为0,同一个TP通信域中各卡的`tp_rank_id`不重复。无TP域通信时,传0即可。135+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当有TP域通信时,取值范围\[0, 1\],默认为0,同一个TP通信域中各卡的`tp_rank_id`不重复。无TP域通信时,传0即可。
138 136 
139-- **expert\_shard\_type** (`int`):可选参数,表示共享专家卡排布类型。137+- **expert\_shard\_type** (`int`):可选参数,表示共享专家卡排布类型。
140- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。138+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。
141- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。139+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前仅支持0,表示共享专家卡排在MoE专家卡前面。
142 140 
143-- **shared\_expert\_num** (`int`):可选参数,表示共享专家数量,一个共享专家可以复制部署到多个卡上。141+- **shared\_expert\_num** (`int`):可选参数,表示共享专家数量,一个共享专家可以复制部署到多个卡上。
144- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。142+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:暂不支持该参数,使用默认值即可。
145- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, 4\],0表示无共享专家,默认值为1。143+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, 4\],0表示无共享专家,默认值为1。
146 144 
147-- **shared\_expert\_rank\_num** (`int`):可选参数,表示共享专家卡数量。145+- **shared\_expert\_rank\_num** (`int`):可选参数,表示共享专家卡数量。
148- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家,使用默认值即可。146+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持共享专家,使用默认值即可。
149- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, ep\_world\_size\)。取0表示无共享专家,不取0需满足shared\_expert\_rank\_num%shared\_expert\_num=0。147+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围\[0, ep\_world\_size\)。取0表示无共享专家,不取0需满足shared\_expert\_rank\_num%shared\_expert\_num=0。
150 148 
151-- **quant\_mode** (`int`):可选参数,表示量化模式。支持取值:0表示非量化(默认),2表示动态量化。当`quant_mode`为2,`dynamic_scales`不为None;当`quant_mode`为0,`dynamic_scales`为None。149+- **quant\_mode** (`int`):可选参数,表示量化模式。支持取值:0表示非量化(默认),2表示动态量化。当`quant_mode`为2,`dynamic_scales`不为None;当`quant_mode`为0,`dynamic_scales`为None。
152-- **global\_bs** (`int`):可选参数,表示EP域全局的batch size大小。当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。150+- **global\_bs** (`int`):可选参数,表示EP域全局的batch size大小。当每个rank的BS不同时,支持传入max\_bs\*ep\_world\_size,其中max\_bs表示单rank BS最大值;当每个rank的BS相同时,支持取值0或BS\*ep\_world\_size。
153 151 
154-- **expert\_token\_nums\_type** (`int`):可选参数,表示输出`expert_token_nums`的值类型,取值范围\[0, 1\],0表示每个专家收到token数量的前缀和,1表示每个专家收到的token数量(默认)。152+- **expert\_token\_nums\_type** (`int`):可选参数,表示输出`expert_token_nums`的值类型,取值范围\[0, 1\],0表示每个专家收到token数量的前缀和,1表示每个专家收到的token数量(默认)。
155 153 
156-- **comm\_alg** (`string`):可选参数,表示通信亲和内存布局算法。154+- **comm\_alg** (`string`):可选参数,表示通信亲和内存布局算法。
157- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当前版本支持"","fullmesh","hierarchy"三种输入方式。推荐配置"hierarchy"并搭配25.0.RC1.1及以上版本驱动使用。155+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当前版本支持"","fullmesh","hierarchy"三种输入方式。推荐配置"hierarchy"并搭配25.0.RC1.1及以上版本驱动使用。
158 - "": 配置HCCL\_INTRA\_PCIE\_ENABLE=1和HCCL\_INTRA\_ROCE\_ENABLE=0时,调用"hierarchy"算法,否则调用"fullmesh"算法。不推荐使用该方式。156 - "": 配置HCCL\_INTRA\_PCIE\_ENABLE=1和HCCL\_INTRA\_ROCE\_ENABLE=0时,调用"hierarchy"算法,否则调用"fullmesh"算法。不推荐使用该方式。
159 - "fullmesh": token数据直接通过RDMA方式发往topk个目标专家所在的卡。157 - "fullmesh": token数据直接通过RDMA方式发往topk个目标专家所在的卡。
160 - "hierarchy": token数据经过跨机、机内两次发送,仅不同server同号卡之间使用RDMA通信,server内使用HCCS通信。158 - "hierarchy": token数据经过跨机、机内两次发送,仅不同server同号卡之间使用RDMA通信,server内使用HCCS通信。
161- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前版本支持"","fullmesh_v1","fullmesh_v2","hierarchy"四种输入方式。159+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:当前版本支持"","fullmesh_v1","fullmesh_v2","hierarchy"四种输入方式。
162 - "":默认值,使能fullmesh_v1模板;160 - "":默认值,使能fullmesh_v1模板;
163 - "fullmesh_v1":使能fullmesh_v1模板;161 - "fullmesh_v1":使能fullmesh_v1模板;
164 - "fullmesh_v2":使能fullmesh_v2模板,其中fullmesh_v2模板仅在tp\_world\_size取值为1时生效。162 - "fullmesh_v2":使能fullmesh_v2模板,其中fullmesh_v2模板仅在tp\_world\_size取值为1时生效。
165 - "hierarchy": token数据经过机内、跨机两次发送,先在server内将同一个token数据汇总求和,再跨机发送,以减少跨机数据量。模板仅支持`tp_world_size`为1、共享专家为0的场景,且不支持二维mask、特殊专家、动态缩容、性能打点场景。163 - "hierarchy": token数据经过机内、跨机两次发送,先在server内将同一个token数据汇总求和,再跨机发送,以减少跨机数据量。模板仅支持`tp_world_size`为1、共享专家为0的场景,且不支持二维mask、特殊专家、动态缩容、性能打点场景。
166 164 
167-- **zero\_expert\_num** (`int`):可选参数,表示零专家的数量。165+- **zero\_expert\_num** (`int`):可选参数,表示零专家的数量。
168- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:166+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
169 - `comm_alg`设置为"fullmesh"时,取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。参数为非0时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。167 - `comm_alg`设置为"fullmesh"时,取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。参数为非0时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。
170 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值0即可。168 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值0即可。
171- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。169+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的零专家的ID值是\[moe\_expert\_num, moe\_expert\_num+zero\_expert\_num\)。
172 170 
173-- **copy\_expert\_num** (`int`):可选参数,表示拷贝专家的数量。171+- **copy\_expert\_num** (`int`):可选参数,表示拷贝专家的数量。
174- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:172+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
175 - `comm_alg`设置为"fullmesh"时,取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。参数为非0时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。173 - `comm_alg`设置为"fullmesh"时,取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。参数为非0时属于零计算专家特性,此特性尚在实验阶段,请谨慎使用。
176 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值0即可。174 - `comm_alg`设置为"hierarchy"时,当前版本不支持,传默认值0即可。
177- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。175+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的拷贝专家的ID值是\[moe\_expert\_num+zero\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num\)。
178 176 
179-- **const\_expert\_num** (`int`):可选参数,表示常量专家的数量。 177+- **const\_expert\_num** (`int`):可选参数,表示常量专家的数量。
180- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当前版本不支持,传0即可。178+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:当前版本不支持,传0即可。
181- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的常量专家的ID值是\[moe\_expert\_num+zero\_expert\_num+copy\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num+const\_expert\_num\)。179+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围[0, MAX_INT32),MAX_INT32 = 2^31 - 1,合法的常量专家的ID值是\[moe\_expert\_num+zero\_expert\_num+copy\_expert\_num, moe\_expert\_num+zero\_expert\_num+copy\_expert\_num+const\_expert\_num\)。
182 180 
183## 输出说明<a name="zh-cn_topic_0000002203575833_section22231435517"></a>181## 输出说明<a name="zh-cn_topic_0000002203575833_section22231435517"></a>
184 182 
185-- **expand\_x** (`Tensor`):表示本卡收到的token数据,要求为2维张量,shape为\(max\(tp\_world\_size, 1\) \*A, H\),A表示在EP通信域可能收到的最大token数,数据类型支持`bfloat16`、`float16`、`int8`。量化时类型为`int8`,非量化时与`x`数据类型保持一致。数据格式为$ND$,支持非连续的Tensor。183+- **expand\_x** (`Tensor`):表示本卡收到的token数据,要求为2维张量,shape为\(max\(tp\_world\_size, 1\) \*A, H\),A表示在EP通信域可能收到的最大token数,数据类型支持`bfloat16`、`float16`、`int8`。量化时类型为`int8`,非量化时与`x`数据类型保持一致。数据格式为$ND$,支持非连续的Tensor。
186-- **dynamic\_scales** (`Tensor`):表示计算得到的动态量化参数。当`quant_mode`不为0时才有该输出,要求为1维张量,shape为\(A,\),数据类型支持`float32`,数据格式支持$ND$,支持非连续的Tensor。184+- **dynamic\_scales** (`Tensor`):表示计算得到的动态量化参数。当`quant_mode`不为0时才有该输出,要求为1维张量,shape为\(A,\),数据类型支持`float32`,数据格式支持$ND$,支持非连续的Tensor。
187-- **assist\_info\_for\_combine** (`Tensor`):表示给同一专家发送的token个数,要求是一个1维张量,shape为\(A \* 128, \)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)的`assist_info_for_combine`输入。185+- **assist\_info\_for\_combine** (`Tensor`):表示给同一专家发送的token个数,要求是一个1维张量,shape为\(A \* 128, \)。数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)的`assist_info_for_combine`输入。
188 186 
189-- **expert\_token\_nums** (`Tensor`):本卡每个专家实际收到的token数量,要求为1维张量,shape为\(local\_expert\_num,\),数据类型`int64`,数据格式支持$ND$,支持非连续的Tensor。187+- **expert\_token\_nums** (`Tensor`):本卡每个专家实际收到的token数量,要求为1维张量,shape为\(local\_expert\_num,\),数据类型`int64`,数据格式支持$ND$,支持非连续的Tensor。
190-- **ep\_recv\_counts** (`Tensor`):表示EP通信域各卡收到的token数(token数以前缀和的形式表示),要求为1维张量,数据类型`int32`,数据格式支持$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)的`ep_send_counts`输入。188+- **ep\_recv\_counts** (`Tensor`):表示EP通信域各卡收到的token数(token数以前缀和的形式表示),要求为1维张量,数据类型`int32`,数据格式支持$ND$,支持非连续的Tensor。对应[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)的`ep_send_counts`输入。
191- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(moe\_expert\_num+2\*global\_bs\*K\*server\_num, \),前`moe_expert_num`个数表示在EP通信域内,该卡上每个专家收到来自其他各卡的token数(以前缀和的形式表示),2\*global\_bs\*K\*server\_num用于存储机间和机内通信前,combine可提前做reduce操作的token个数和通信区偏移量,`global_bs`传入0时此处按照bs\*ep\_world\_size计算。189+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求shape为\(moe\_expert\_num+2\*global\_bs\*K\*server\_num, \),前`moe_expert_num`个数表示在EP通信域内,该卡上每个专家收到来自其他各卡的token数(以前缀和的形式表示),2\*global\_bs\*K\*server\_num用于存储机间和机内通信前,combine可提前做reduce操作的token个数和通信区偏移量,`global_bs`传入0时此处按照bs\*ep\_world\_size计算。
192- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(ep\_world\_size\*max\(tp\_world\_size, 1\)\*local\_expert\_num, \)。190+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:要求shape为\(ep\_world\_size\*max\(tp\_world\_size, 1\)\*local\_expert\_num, \)。
193 191 
194-- **tp\_recv\_counts** (`Tensor`):表示TP通信域各卡收到的token数量。对应[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)的`tp_send_counts`输入。192+- **tp\_recv\_counts** (`Tensor`):表示TP通信域各卡收到的token数量。对应[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)的`tp_send_counts`输入。
195- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP通信域,暂无该输出,193+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持TP通信域,暂无该输出,
196- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求是一个1D Tensor,shape为\(tp\_world\_size, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。194+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:支持TP通信域,要求是一个1D Tensor,shape为\(tp\_world\_size, \),数据类型支持`int32`,数据格式为$ND$,支持非连续的Tensor。
197 195 
198-- **expand\_scales** (`Tensor`):表示`expert_scales`与`x`一起进行alltoallv之后的输出。196+- **expand\_scales** (`Tensor`):表示`expert_scales`与`x`一起进行alltoallv之后的输出。
199- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求是一个1维张量,shape为\(A, \),数据类型支持`float32`,数据格式要求为$ND$,支持非连续的Tensor。197+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:要求是一个1维张量,shape为\(A, \),数据类型支持`float32`,数据格式要求为$ND$,支持非连续的Tensor。
200- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`comm_alg`设置为"hierarchy"时,要求是一个1维张量,shape为\(A, \),数据类型支持`float32`,数据格式要求为$ND$,支持非连续的Tensor。`comm_alg`设置为"","fullmesh_v1","fullmesh_v2"时,暂不支持该输出,返回None。198+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`comm_alg`设置为"hierarchy"时,要求是一个1维张量,shape为\(A, \),数据类型支持`float32`,数据格式要求为$ND$,支持非连续的Tensor。`comm_alg`设置为"","fullmesh_v1","fullmesh_v2"时,暂不支持该输出,返回None。
201 199 
202## 约束说明<a name="zh-cn_topic_0000002203575833_section12345537164214"></a>200## 约束说明<a name="zh-cn_topic_0000002203575833_section12345537164214"></a>
203 201 
204-- 该接口支持推理场景下使用。202+- 该接口支持推理场景下使用。
205-- 该接口支持静态图模式,`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`必须配套使用。203+- 该接口支持静态图模式,`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`必须配套使用。
206-- 在不同产品型号、不同通信算法或不同版本中,`npu_moe_distribute_dispatch_v2`的Tensor输出`assist_info_for_combine`、`ep_recv_counts`、`tp_recv_counts`、`expand_scales`中的元素值可能不同,使用时直接将上述Tensor传给`npu_moe_distribute_combine_v2`对应参数即可,模型其他业务逻辑不应对其存在依赖。204+- 在不同产品型号、不同通信算法或不同版本中,`npu_moe_distribute_dispatch_v2`的Tensor输出`assist_info_for_combine`、`ep_recv_counts`、`tp_recv_counts`、`expand_scales`中的元素值可能不同,使用时直接将上述Tensor传给`npu_moe_distribute_combine_v2`对应参数即可,模型其他业务逻辑不应对其存在依赖。
207-- 调用接口过程中使用的`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`shared_expert_num`、`shared_expert_rank_num`、`global_bs`参数取值所有卡需保持一致,`group_ep`、`ep_world_size`、`group_tp`、`tp_world_size`、`expert_shard_type`、`global_bs`网络不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)对应参数也保持一致。205+- 调用接口过程中使用的`group_ep`、`ep_world_size`、`moe_expert_num`、`group_tp`、`tp_world_size`、`expert_shard_type`、`shared_expert_num`、`shared_expert_rank_num`、`global_bs`参数取值所有卡需保持一致,`group_ep`、`ep_world_size`、`group_tp`、`tp_world_size`、`expert_shard_type`、`global_bs`网络不同层中也需保持一致,且和[torch\_npu.npu\_moe\_distribute\_combine\_v2](torch_npu-npu_moe_distribute_combine_v2.md)对应参数也保持一致。
208-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。206+- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。
209-- moe_expert_num + zero_expert_num + copy_expert_num + const_expert_num < MAX_INT32。207+- moe_expert_num + zero_expert_num + copy_expert_num + const_expert_num < MAX_INT32。
210-- 参数里Shape使用的变量如下:208+- 参数里Shape使用的变量如下:
211- - A:表示本卡接收的最大token数量,取值范围如下209+ - A:表示本卡接收的最大token数量,取值范围如下
212- - 对于共享专家,要满足A=BS\*shared\_expert\_num/shared\_expert\_rank\_num。210+ - 对于共享专家,要满足A=BS\*shared\_expert\_num/shared\_expert\_rank\_num。
213- - 对于MoE专家,当`global_bs`为0时,要满足A\>=BS\*ep\_world\_size\*min\(local\_expert\_num, K\);当`global_bs`不为0时,要满足A\>=global\_bs\* min\(local\_expert\_num, K\)。211+ - 对于MoE专家,当`global_bs`为0时,要满足A\>=BS\*ep\_world\_size\*min\(local\_expert\_num, K\);当`global_bs`不为0时,要满足A\>=global\_bs\* min\(local\_expert\_num, K\)。
214 212 
215- - H:表示hidden size隐藏层大小。213+ - H:表示hidden size隐藏层大小。
216- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`H`的取值范围如下所示。214+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`H`的取值范围如下所示。
217 - `comm_alg`设置为"fullmesh"时,`H`的取值范围\(0, 7168\],且保证是32的整数倍。215 - `comm_alg`设置为"fullmesh"时,`H`的取值范围\(0, 7168\],且保证是32的整数倍。
218 - `comm_alg`设置为"hierarchy"且驱动版本不低于25.0.RC1.1时,`H`的取值范围\(0, 10 * 1024\],且保证是32的整数倍。216 - `comm_alg`设置为"hierarchy"且驱动版本不低于25.0.RC1.1时,`H`的取值范围\(0, 10 * 1024\],且保证是32的整数倍。
219- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值为\[1024, 8192\]。217+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值为\[1024, 8192\]。
220 218 
221- - BS:表示batch sequence size,即本卡最终输出的token数量。219+ - BS:表示batch sequence size,即本卡最终输出的token数量。
222- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`BS`的取值范围如下所示。220+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:`BS`的取值范围如下所示。
223 - `comm_alg`设置为"fullmesh"时,`BS`的取值范围为0<BS≤256。221 - `comm_alg`设置为"fullmesh"时,`BS`的取值范围为0<BS≤256。
224 - `comm_alg`设置为"hierarchy"且Ascend HDK版本不低于25.0.RC1.1时,`BS`的取值范围为0<BS≤512。222 - `comm_alg`设置为"hierarchy"且Ascend HDK版本不低于25.0.RC1.1时,`BS`的取值范围为0<BS≤512。
225- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`BS`的取值范围如下所示。223+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:`BS`的取值范围如下所示。
226 - `comm_alg`设置为""或"fullmesh_v1"时,`BS`的取值范围为0<BS≤512。224 - `comm_alg`设置为""或"fullmesh_v1"时,`BS`的取值范围为0<BS≤512。
227 - `comm_alg`设置为"fullmesh_v2"或"hierarchy"时,`BS`的取值范围为0<BS≤256。225 - `comm_alg`设置为"fullmesh_v2"或"hierarchy"时,`BS`的取值范围为0<BS≤256。
228 226 
229- - K:表示选取topK个专家,取值范围为0<K≤16,同时满足0 < K ≤ moe\_expert\_num + zero_expert_num + copy_expert_num + const_expert_num,当`comm_alg`为"fullmesh_v2"时,取值范围为0<K≤12。227+ - K:表示选取topK个专家,取值范围为0<K≤16,同时满足0 < K ≤ moe\_expert\_num + zero_expert_num + copy_expert_num + const_expert_num,当`comm_alg`为"fullmesh_v2"时,取值范围为0<K≤12。
230 228 
231- - server\_num:表示服务器的节点数,取值只支持2、4、8。229+ - server\_num:表示服务器的节点数,取值只支持2、4、8。
232- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:仅该场景的shape使用了该变量。230+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:仅该场景的shape使用了该变量。
233 231 
234- - local\_expert\_num:表示本卡专家数量。232+ - local\_expert\_num:表示本卡专家数量。
235- - 对于共享专家卡,local\_expert\_num为1。233+ - 对于共享专家卡,local\_expert\_num为1。
236- - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num),当local_expert_num大于1时,不支持TP域通信。234+ - 对于MoE专家卡,local\_expert\_num=moe\_expert\_num/\(ep\_world\_size-shared\_expert\_rank\_num),当local_expert_num大于1时,不支持TP域通信。
237- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:应满足0 < local\_expert\_num * ep\_world\_size ≤ 2048。235+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:应满足0 < local\_expert\_num * ep\_world\_size ≤ 2048。
238 236 
239-- HCCL通信域缓存区大小:237+- HCCL通信域缓存区大小:
240 238 
241 调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。239 调用本接口前需检查通信域缓存区大小取值是否合理,单位MB,不配置时默认为200MB。
242- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:240+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:
243 该场景支持通过环境变量HCCL\_BUFFSIZE配置。241 该场景支持通过环境变量HCCL\_BUFFSIZE配置。
244 - `comm_alg`配置为"": 依照HCCL\_INTRA\_PCIE\_ENABLE和HCCL\_INTRA\_ROCE\_ENABLE配置选择"fullmesh"或"hierarchy"公式。242 - `comm_alg`配置为"": 依照HCCL\_INTRA\_PCIE\_ENABLE和HCCL\_INTRA\_ROCE\_ENABLE配置选择"fullmesh"或"hierarchy"公式。
245 - `comm_alg`配置为"fullmesh": 设置大小要求\>=2\*\(BS\*ep\_world\_size\*min\(local\_expert\_num, K\)\*H\*sizeof\(uint16\)+2MB\)。243 - `comm_alg`配置为"fullmesh": 设置大小要求\>=2\*\(BS\*ep\_world\_size\*min\(local\_expert\_num, K\)\*H\*sizeof\(uint16\)+2MB\)。
246 - `comm_alg`配置为"hierarchy": 设置大小要求 \>= \(moe\_expert\_num + ep\_world\_size / 4\) \* Align512\(max_bs \* \(H \* sizeof\(dtype_x\) + 4 \* Align8\(K\) \* sizeof\(uint32\)\)\) \* 1B + 8MB,其中Align512\(x\) = \(\(x+512-1\)/512\)\*512,Align8\(x\) = \(\(x+8-1\)/8\)\*8。244 - `comm_alg`配置为"hierarchy": 设置大小要求 \>= \(moe\_expert\_num + ep\_world\_size / 4\) \* Align512\(max_bs \* \(H \* sizeof\(dtype_x\) + 4 \* Align8\(K\) \* sizeof\(uint32\)\)\) \* 1B + 8MB,其中Align512\(x\) = \(\(x+512-1\)/512\)\*512,Align8\(x\) = \(\(x+8-1\)/8\)\*8。
247 245 
248- - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:246+ - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:
249 该场景不仅支持通过环境变量HCCL\_BUFFSIZE配置,还支持通过hccl_buffer_size配置(参考《[PyTorch训练模型迁移调优](https://hiascend.com/document/redirect/canncommercial-ptmigr)》中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。247 该场景不仅支持通过环境变量HCCL\_BUFFSIZE配置,还支持通过hccl_buffer_size配置(参考《[PyTorch训练模型迁移调优](https://hiascend.com/document/redirect/canncommercial-ptmigr)》中“性能调优>性能调优方法>通信优化>优化方法>hccl_buffer_size”章节)。
250 - ep通信域内,comm\_alg配置为"fullmesh_v1"或"": 设置大小要求 \>= 2 \* \(local\_expert\_num \* max\_bs \* ep\_world\_size \* Align512\(Align32\(2 \* H\) + 64\) + \(K + shared\_expert\_num\) \* max\_bs \* Align512\(2 \* H\)\)。248 - ep通信域内,comm\_alg配置为"fullmesh_v1"或"": 设置大小要求 \>= 2 \* \(local\_expert\_num \* max\_bs \* ep\_world\_size \* Align512\(Align32\(2 \* H\) + 64\) + \(K + shared\_expert\_num\) \* max\_bs \* Align512\(2 \* H\)\)。
251 - ep通信域内,comm\_alg配置为"fullmesh_v2": 设置大小要求 \>= 2 \* \(local\_expert\_num \* max\_bs \* ep\_world\_size \* 480Align512\(Align32\(2 \* H\) + 64\) + \(K + shared\_expert\_num\) \* max\_bs \* Align512\(2 \* H\)\)。249 - ep通信域内,comm\_alg配置为"fullmesh_v2": 设置大小要求 \>= 2 \* \(local\_expert\_num \* max\_bs \* ep\_world\_size \* 480Align512\(Align32\(2 \* H\) + 64\) + \(K + shared\_expert\_num\) \* max\_bs \* Align512\(2 \* H\)\)。
@@ -253,27 +251,27 @@ torch_npu.npu_moe_distribute_dispatch_v2(x, expert_ids, group_ep, ep_world_size,
253 - tp通信域内:设置大小要求 \>= (A \* Align512(Align32(h \* 2) + 44) + A \* Align512(h \* 2)) \* 2。251 - tp通信域内:设置大小要求 \>= (A \* Align512(Align32(h \* 2) + 44) + A \* Align512(h \* 2)) \* 2。
254 - 其中 480Align512(x) = ((x+480-1)/480)\*512,Align512(x) = ((x+512-1)/512)\*512,Align32(x) = ((x+32-1)/32)\*32。252 - 其中 480Align512(x) = ((x+480-1)/480)\*512,Align512(x) = ((x+512-1)/512)\*512,Align32(x) = ((x+32-1)/32)\*32。
255 253 
256-- HCCL_INTRA_PCIE_ENABLE和HCCL_INTRA_ROCE_ENABLE:254+- HCCL_INTRA_PCIE_ENABLE和HCCL_INTRA_ROCE_ENABLE:
257- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:该环境变量不再推荐使用,建议comm\_alg配置"hierarchy"。255+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:该环境变量不再推荐使用,建议comm\_alg配置"hierarchy"。
258 256 
259-- 本文公式中的“/”表示整除。257+- 本文公式中的“/”表示整除。
260 258 
261-- 通信域使用约束:259+- 通信域使用约束:
262 260 
263- - 一个模型中的`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`算子仅支持相同EP通信域,且该通信域中不允许有其他算子。261+ - 一个模型中的`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`算子仅支持相同EP通信域,且该通信域中不允许有其他算子。
264 262 
265- - 一个模型中的`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。263+ - 一个模型中的`npu_moe_distribute_dispatch_v2`和`npu_moe_distribute_combine_v2`算子仅支持相同TP通信域或都不支持TP通信域,有TP通信域时该通信域中不允许有其他算子。
266 264 
267-- 组网约束:265+- 组网约束:
268- - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:多机场景仅支持交换机组网,不支持双机直连组网。266+ - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:多机场景仅支持交换机组网,不支持双机直连组网。
269 267 
270-- 版本配套约束:268+- 版本配套约束:
271 269 
272 静态图模式下,从Ascend Extension for PyTorch 8.0.0版本开始,Ascend Extension for PyTorch框架会对静态图中最后一个节点输出结果做Meta推导与inferShape推导的结果强校验。当图中只有一个Dispatch\_v2算子,若CANN版本落后于Ascend Extension for PyTorch版本,会出现Shape不匹配报错,建议用户升级CANN版本,详细的版本配套关系参见《[Ascend Extension for PyTorch 版本说明](https://gitcode.com/Ascend/pytorch/blob/v2.7.1-7.3.0/docs/zh/release_notes/release_notes.md)》中“相关产品版本配套说明”。270 静态图模式下,从Ascend Extension for PyTorch 8.0.0版本开始,Ascend Extension for PyTorch框架会对静态图中最后一个节点输出结果做Meta推导与inferShape推导的结果强校验。当图中只有一个Dispatch\_v2算子,若CANN版本落后于Ascend Extension for PyTorch版本,会出现Shape不匹配报错,建议用户升级CANN版本,详细的版本配套关系参见《[Ascend Extension for PyTorch 版本说明](https://gitcode.com/Ascend/pytorch/blob/v2.7.1-7.3.0/docs/zh/release_notes/release_notes.md)》中“相关产品版本配套说明”。
273 271 
274## 调用示例<a name="zh-cn_topic_0000002203575833_section14459801435"></a>272## 调用示例<a name="zh-cn_topic_0000002203575833_section14459801435"></a>
275 273 
276-- 单算子模式调用274+- 单算子模式调用
277 275 
278 ```python276 ```python
279 import os277 import os
@@ -514,7 +512,7 @@ torch_npu.npu_moe_distribute_dispatch_v2(x, expert_ids, group_ep, ep_world_size,
514 print("run npu success.")512 print("run npu success.")
515 ```513 ```
516 514 
517-- 图模式调用515+- 图模式调用
518 516 
519 ```python517 ```python
520 # 仅支持静态图518 # 仅支持静态图
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_finalize_routing.md+15-17
@@ -93,18 +93,20 @@
93 93 
94## 函数原型94## 函数原型
95 95 
96-```96+```python
97torch_npu.npu_moe_finalize_routing(expanded_permuted_rows, skip1, skip2, bias, scales, expanded_src_to_dst_row, export_for_source_row, drop_pad_mode=0) -> Tensor97torch_npu.npu_moe_finalize_routing(expanded_permuted_rows, skip1, skip2, bias, scales, expanded_src_to_dst_row, export_for_source_row, drop_pad_mode=0) -> Tensor
98```98```
99 99 
100## 参数说明100## 参数说明
101+>
101> [!NOTE] 102> [!NOTE]
102> shape中的符号说明:103> shape中的符号说明:
103-> - $NUM\_ROWS$:为行数。104+>
104-> - $K$:表示从总的专家$E$中选出$K$个专家105+> - $NUM\_ROWS$:为行数
105-> - $H$:表示token序列长度,为列数106+> - $K$:表示从总的专家$E$中选出$K$专家
106-> - $E$: 表示专家数$E$需要大于等于$K$107+> - $H$表示每个token序列长度为列数
107-> - $C$: 表示专家处理token量的能力阈值108+> - $E$: 表示专家数,$E$需要大于等于$K$。
109+> - $C$: 表示专家处理token数量的能力阈值
108 110 
109- **expanded_permuted_rows** (`Tensor`):必选参数,对应公式中的$expandPermutedRows$,经过专家处理过的结果,要求为一个2维张量,数据类型支持`float16``bfloat16``float32`,数据格式要求为$ND$。`drop_pad_mode`参数为0或2时,shape为$(NUM\_ROWS * K, H)$,`drop_pad_mode`参数为1或3时,shape为$(E, C, H)$。111- **expanded_permuted_rows** (`Tensor`):必选参数,对应公式中的$expandPermutedRows$,经过专家处理过的结果,要求为一个2维张量,数据类型支持`float16``bfloat16``float32`,数据格式要求为$ND$。`drop_pad_mode`参数为0或2时,shape为$(NUM\_ROWS * K, H)$,`drop_pad_mode`参数为1或3时,shape为$(E, C, H)$。
110- **skip1** (`Tensor`):必选参数,允许为None,对应公式中的$skip1$,求和的输入参数1,要求为一个2维张量,数据类型要求与`expanded_permuted_rows`一致,shape要求与输出`out`的shape一致。112- **skip1** (`Tensor`):必选参数,允许为None,对应公式中的$skip1$,求和的输入参数1,要求为一个2维张量,数据类型要求与`expanded_permuted_rows`一致,shape要求与输出`out`的shape一致。
@@ -114,13 +116,13 @@ torch_npu.npu_moe_finalize_routing(expanded_permuted_rows, skip1, skip2, bias, s
114- **expanded_src_to_dst_row** (`Tensor`):必选参数,对应公式中的$expandedSrcToDstRow$,保存每个专家处理结果的索引,要求为一个1维张量,数据类型支持`int32`。shape支持$(NUM\_ROWS * K)$,`drop_pad_mode`参数为0或2时,Tensor的取值范围是$[0, NUM\_ROWS * K-1]$,`drop_pad_mode`参数为1或3时,Tensor的取值范围是$[-1, E*C - 1]$。116- **expanded_src_to_dst_row** (`Tensor`):必选参数,对应公式中的$expandedSrcToDstRow$,保存每个专家处理结果的索引,要求为一个1维张量,数据类型支持`int32`。shape支持$(NUM\_ROWS * K)$,`drop_pad_mode`参数为0或2时,Tensor的取值范围是$[0, NUM\_ROWS * K-1]$,`drop_pad_mode`参数为1或3时,Tensor的取值范围是$[-1, E*C - 1]$。
115- **export_for_source_row** (`Tensor`):必选参数,允许为None,公式中的$exportForSourceRow$,每行处理的专家号,要求为一个2维张量,数据类型支持`int32`。shape支持$(NUM\_ROWS,K)$,Tensor的取值范围是[0,E-1]。117- **export_for_source_row** (`Tensor`):必选参数,允许为None,公式中的$exportForSourceRow$,每行处理的专家号,要求为一个2维张量,数据类型支持`int32`。shape支持$(NUM\_ROWS,K)$,Tensor的取值范围是[0,E-1]。
116- **drop_pad_mode** (`int`):可选参数,表示是否为丢弃模式(丢弃模式即drop pad模式,非丢弃模式即drop less模式)和`expanded_src_to_dst_row`的排列方式(行排列或列排列),取值范围为[0, 3],默认值为`0`118- **drop_pad_mode** (`int`):可选参数,表示是否为丢弃模式(丢弃模式即drop pad模式,非丢弃模式即drop less模式)和`expanded_src_to_dst_row`的排列方式(行排列或列排列),取值范围为[0, 3],默认值为`0`
117- - 0表示非丢弃模式,`expanded_src_to_dst_row`按列排列。119+ - 0表示非丢弃模式,`expanded_src_to_dst_row`按列排列。
118- - 1表示丢弃模式,`expanded_src_to_dst_row`按列排列。120+ - 1表示丢弃模式,`expanded_src_to_dst_row`按列排列。
119- - 2表示非丢弃模式场景,`expanded_src_to_dst_row`按行排列。121+ - 2表示非丢弃模式场景,`expanded_src_to_dst_row`按行排列。
120- - 3表示丢弃场景,`expanded_src_to_dst_row`按行排列。122+ - 3表示丢弃场景,`expanded_src_to_dst_row`按行排列。
121- 
122 123 
123## 返回值说明124## 返回值说明
125+ 
124`Tensor`126`Tensor`
125 127 
126输出参数`out`,代表最后MoE FFN合并的输出结果。数据维度支持2维,shape支持$(NUM\_ROWS, H)$,dtype与expanded_permuted_rows保持一致。128输出参数`out`,代表最后MoE FFN合并的输出结果。数据维度支持2维,shape支持$(NUM\_ROWS, H)$,dtype与expanded_permuted_rows保持一致。
@@ -134,6 +136,7 @@ torch_npu.npu_moe_finalize_routing(expanded_permuted_rows, skip1, skip2, bias, s
134 136 
135- 单算子模式调用137- 单算子模式调用
136 - drop less模式调用示例138 - drop less模式调用示例
139+ 
137 ```python140 ```python
138 >>> import torch141 >>> import torch
139 >>> import torch_npu142 >>> import torch_npu
@@ -159,6 +162,7 @@ torch_npu.npu_moe_finalize_routing(expanded_permuted_rows, skip1, skip2, bias, s
159 ```162 ```
160 163
161 - drop pad模式调用示例164 - drop pad模式调用示例
165+ 
162 ```python166 ```python
163 >>> import torch167 >>> import torch
164 >>> import torch_npu168 >>> import torch_npu
@@ -184,7 +188,6 @@ torch_npu.npu_moe_finalize_routing(expanded_permuted_rows, skip1, skip2, bias, s
184 torch.float32188 torch.float32
185 ```189 ```
186 190 
187- 
188- 图模式调用191- 图模式调用
189 192 
190 ```python193 ```python
@@ -233,8 +236,3 @@ torch_npu.npu_moe_finalize_routing(expanded_permuted_rows, skip1, skip2, bias, s
233 # 执行上述代码的输出类似如下236 # 执行上述代码的输出类似如下
234 torch.Size([50, 10]) torch.float32237 torch.Size([50, 10]) torch.float32
235 ```238 ```
236- 
237- 
238- 
239- 
240- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_gating_top_k.md+24-29
@@ -9,8 +9,8 @@
9 9 
10## 功能说明10## 功能说明
11 11 
12-- API功能:MoE计算中,对输入x做Sigmoid/SoftMax计算,对计算结果分组进行排序,最后根据分组排序的结果选取前k个专家。12+- API功能:MoE计算中,对输入x做Sigmoid/SoftMax计算,对计算结果分组进行排序,最后根据分组排序的结果选取前k个专家。
13-- 计算公式:13+- 计算公式:
14 14 
15 对输入做sigmoid(`bias`可选):15 对输入做sigmoid(`bias`可选):
16 -`norm_type`为1时:16 -`norm_type`为1时:
@@ -39,7 +39,7 @@
39 39 
40 ![](../../figures/zh-cn_formulaimage_0000002219173660.png)40 ![](../../figures/zh-cn_formulaimage_0000002219173660.png)
41 41 
42-- 等价计算逻辑:42+- 等价计算逻辑:
43 43 
44 ```python44 ```python
45 import torch45 import torch
@@ -119,47 +119,47 @@
119 119 
120## 函数原型120## 函数原型
121 121 
122-```122+```python
123npu_moe_gating_top_k(x, k, *, bias=None, k_group=1, group_count=1, group_select_mode=0, renorm=0, norm_type=0, out_flag=False, routed_scaling_factor=1.0, eps=1e-20) -> (Tensor, Tensor, Tensor)123npu_moe_gating_top_k(x, k, *, bias=None, k_group=1, group_count=1, group_select_mode=0, renorm=0, norm_type=0, out_flag=False, routed_scaling_factor=1.0, eps=1e-20) -> (Tensor, Tensor, Tensor)
124```124```
125 125 
126## 参数说明126## 参数说明
127 127 
128-- **x**(`Tensor`):必选参数,表示待计算的输入。要求是一个2D的Tensor,数据类型支持`float16`、`bfloat16`、`float32`,数据格式要求为ND。支持非连续Tensor。最后一维的大小(即专家数)要求不大于`2048`。128+- **x**(`Tensor`):必选参数,表示待计算的输入。要求是一个2D的Tensor,数据类型支持`float16`、`bfloat16`、`float32`,数据格式要求为ND。支持非连续Tensor。最后一维的大小(即专家数)要求不大于`2048`。
129 129 
130-- **k**(`int`):必选参数,表示每个token最终筛选得到的专家个数,数据类型为`int64`。要求`1 <= k <= x_shape[-1] / group_count * k_group`。130+- **k**(`int`):必选参数,表示每个token最终筛选得到的专家个数,数据类型为`int64`。要求`1 <= k <= x_shape[-1] / group_count * k_group`。
131 131 
132-- <strong>*</strong>:代表其之前的变量是位置相关,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。132+- <strong>*</strong>:代表其之前的变量是位置相关,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
133 133
134-- **bias**(`Tensor`):可选参数,表示与输入`x`进行计算的bias值。要求是1D的Tensor,要求shape值与`x`的最后一维相等。数据类型支持`float16`、`bfloat16`、`float32`,数据类型需要与`x`保持一致,数据格式要求为ND。支持非连续`Tensor`。134+- **bias**(`Tensor`):可选参数,表示与输入`x`进行计算的bias值。要求是1D的Tensor,要求shape值与`x`的最后一维相等。数据类型支持`float16`、`bfloat16`、`float32`,数据类型需要与`x`保持一致,数据格式要求为ND。支持非连续`Tensor`。
135 135 
136-- **k_group**(`int`):可选参数,表示每个token组筛选过程中,选出的专家组个数,数据类型为`int64`,默认值为`1`。要求`1 <= k_group <= group_count`,并且`k_group * x_shape[-1] / group_count`的值要大于等于`k`。136+- **k_group**(`int`):可选参数,表示每个token组筛选过程中,选出的专家组个数,数据类型为`int64`,默认值为`1`。要求`1 <= k_group <= group_count`,并且`k_group * x_shape[-1] / group_count`的值要大于等于`k`。
137 137 
138-- **group_count**(`int`):可选参数,表示将全部专家划分的组数,数据类型为`int64`,默认值为`1`。要求group_count > 0,x_shape[-1]能够被`group_count`整除且整除后的结果大于`2`,并且整除的结果按照32个数对齐后乘`group_count`的结果不大于`2048`。138+- **group_count**(`int`):可选参数,表示将全部专家划分的组数,数据类型为`int64`,默认值为`1`。要求group_count > 0,x_shape[-1]能够被`group_count`整除且整除后的结果大于`2`,并且整除的结果按照32个数对齐后乘`group_count`的结果不大于`2048`。
139 139 
140-- **group_select_mode**(`int`):可选参数,表示一个专家组的总得分计算方式。默认值为`0`,`0`表示组内取最大值,作为专家组得分;`1`表示取组内Top2的专家进行得分累加,作为专家组得分。140+- **group_select_mode**(`int`):可选参数,表示一个专家组的总得分计算方式。默认值为`0`,`0`表示组内取最大值,作为专家组得分;`1`表示取组内Top2的专家进行得分累加,作为专家组得分。
141 141 
142-- **renorm**(`int`):可选参数,表示renorm标记,默认值为`0`,表示先进行norm再进行topk计算。当前仅支持`0`。142+- **renorm**(`int`):可选参数,表示renorm标记,默认值为`0`,表示先进行norm再进行topk计算。当前仅支持`0`。
143-- **norm_type**(`int`):可选参数,表示norm函数类型,`1`表示使用Sigmoid函数,`0`表示Softmax函数。默认值为`0`。143+- **norm_type**(`int`):可选参数,表示norm函数类型,`1`表示使用Sigmoid函数,`0`表示Softmax函数。默认值为`0`。
144 144 
145-- **out_flag**(`bool`):可选参数,是否输出norm函数中间结果。默认值为`False`。145+- **out_flag**(`bool`):可选参数,是否输出norm函数中间结果。默认值为`False`。
146-- **routed_scaling_factor**(`float`):可选参数,表示计算`yOut`使用的`routed_scaling_factor`系数,默认值为`1.0`。146+- **routed_scaling_factor**(`float`):可选参数,表示计算`yOut`使用的`routed_scaling_factor`系数,默认值为`1.0`。
147-- **eps**(`float`):可选参数,表示计算`yOut`使用的`eps`系数,默认值为`1e-20`。147+- **eps**(`float`):可选参数,表示计算`yOut`使用的`eps`系数,默认值为`1e-20`。
148 148 
149## 返回值说明149## 返回值说明
150 150 
151-- **yOut**(`Tensor`):表示对`x`做norm操作和分组排序topk后计算的结果。要求是一个2D的Tensor,数据类型支持`float16`、`bfloat16`、`float32`,数据类型与`x`需要保持一致,数据格式要求为ND,第一维的大小要求与`x`的第一维相同,最后一维的大小与`k`相同。不支持非连续Tensor。151+- **yOut**(`Tensor`):表示对`x`做norm操作和分组排序topk后计算的结果。要求是一个2D的Tensor,数据类型支持`float16`、`bfloat16`、`float32`,数据类型与`x`需要保持一致,数据格式要求为ND,第一维的大小要求与`x`的第一维相同,最后一维的大小与`k`相同。不支持非连续Tensor。
152-- **expertIdxOut**(`Tensor`):表示对`x`做norm操作和分组排序topk后的索引,即专家的序号。shape要求与yOut一致,数据类型支持`int32`,数据格式要求为ND。不支持非连续Tensor。152+- **expertIdxOut**(`Tensor`):表示对`x`做norm操作和分组排序topk后的索引,即专家的序号。shape要求与yOut一致,数据类型支持`int32`,数据格式要求为ND。不支持非连续Tensor。
153-- **normOut**(`Tensor`):表示norm计算的输出结果。shape要求与`x`保持一致,数据类型为`float32`,数据格式要求为ND。不支持非连续Tensor。153+- **normOut**(`Tensor`):表示norm计算的输出结果。shape要求与`x`保持一致,数据类型为`float32`,数据格式要求为ND。不支持非连续Tensor。
154 154 
155## 约束说明155## 约束说明
156 156 
157-- 该接口支持推理场景下使用。157+- 该接口支持推理场景下使用。
158-- 该接口支持图模式。158+- 该接口支持图模式。
159 159 
160## 调用示例160## 调用示例
161 161 
162-- 单算子模式调用162+- 单算子模式调用
163 163 
164 ```python164 ```python
165 import torch165 import torch
@@ -186,7 +186,7 @@ npu_moe_gating_top_k(x, k, *, bias=None, k_group=1, group_count=1, group_select_
186 y_npu, expert_idx_npu, out_npu = torch_npu.npu_moe_gating_top_k(x_tensor, k, bias=bias_tensor, k_group=k_group, group_count=group_count, group_select_mode=group_select_mode, renorm=renorm, norm_type=norm_type, out_flag=out_flag, routed_scaling_factor=routed_scaling_factor, eps=eps)186 y_npu, expert_idx_npu, out_npu = torch_npu.npu_moe_gating_top_k(x_tensor, k, bias=bias_tensor, k_group=k_group, group_count=group_count, group_select_mode=group_select_mode, renorm=renorm, norm_type=norm_type, out_flag=out_flag, routed_scaling_factor=routed_scaling_factor, eps=eps)
187 ```187 ```
188 188 
189-- 图模式调用189+- 图模式调用
190 190 
191 ```python191 ```python
192 # 入图方式192 # 入图方式
@@ -227,9 +227,4 @@ npu_moe_gating_top_k(x, k, *, bias=None, k_group=1, group_count=1, group_select_
227 # 调用MoeGatingTopK算子227 # 调用MoeGatingTopK算子
228 y_npu, expert_idx_npu, out_npu = model(x_tensor, bias_tensor)228 y_npu, expert_idx_npu, out_npu = model(x_tensor, bias_tensor)
229 ```229 ```
230- 230+
231- 
232-
233- 
234- 
235- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_gating_top_k_softmax.md+1-2
@@ -20,7 +20,7 @@ $$
20 20 
21## 函数原型21## 函数原型
22 22 
23-```23+```python
24torch_npu.npu_moe_gating_top_k_softmax(x, finished=None, k=1) -> (Tensor, Tensor, Tensor)24torch_npu.npu_moe_gating_top_k_softmax(x, finished=None, k=1) -> (Tensor, Tensor, Tensor)
25```25```
26 26 
@@ -77,4 +77,3 @@ torch_npu.npu_moe_gating_top_k_softmax(x, finished=None, k=1) -> (Tensor, Tensor
77 res = moe_gating_topk_softmax_model(x, None, 2)77 res = moe_gating_topk_softmax_model(x, None, 2)
78 print(res)78 print(res)
79 ```79 ```
80- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_init_routing.md+2-4
@@ -57,7 +57,7 @@
57 57
58## 函数原型58## 函数原型
59 59 
60-```60+```python
61torch_npu.npu_moe_init_routing(x, row_idx, expert_idx, active_num) -> (Tensor, Tensor, Tensor)61torch_npu.npu_moe_init_routing(x, row_idx, expert_idx, active_num) -> (Tensor, Tensor, Tensor)
62```62```
63 63 
@@ -128,6 +128,4 @@ torch_npu.npu_moe_init_routing(x, row_idx, expert_idx, active_num) -> (Tensor, T
128 print(expanded_row_idx)128 print(expanded_row_idx)
129 print(expanded_expert_idx)129 print(expanded_expert_idx)
130 ```130 ```
131- 131+
132-
133- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_init_routing_v2.md+51-53
@@ -9,7 +9,7 @@
9 9 
10## 功能说明<a name="zh-cn_topic_0000002271534921_section1650913464367"></a>10## 功能说明<a name="zh-cn_topic_0000002271534921_section1650913464367"></a>
11 11 
12-- API功能:MoE(Mixture of Experts)的routing计算,根据[torch_npu.npu_moe_gating_top_k_softmax](torch_npu-npu_moe_gating_top_k_softmax.md)的计算结果做routing处理,支持不量化、动态量化和静态量化模式。12+- API功能:MoE(Mixture of Experts)的routing计算,根据[torch_npu.npu_moe_gating_top_k_softmax](torch_npu-npu_moe_gating_top_k_softmax.md)的计算结果做routing处理,支持不量化、动态量化和静态量化模式。
13- 计算公式: 13- 计算公式:
14 14 
15 1.对输入expertIdx做排序,得出排序后的结果sortedExpertIdx和对应的序号sortedRowIdx:15 1.对输入expertIdx做排序,得出排序后的结果sortedExpertIdx和对应的序号sortedRowIdx:
@@ -40,7 +40,7 @@
40 quantResult=round((x∗scaleOptional)+offsetOptional)40 quantResult=round((x∗scaleOptional)+offsetOptional)
41 $$41 $$
42 42
43- - 动态量化:43+ - 动态量化:
44 - 若不输入scale:44 - 若不输入scale:
45 $$45 $$
46 dynamicQuantScaleOutOptional = row\_max(abs(x)) / 12746 dynamicQuantScaleOutOptional = row\_max(abs(x)) / 127
@@ -83,7 +83,8 @@
83 availableIdxNum = |\{x\in expertIdx| expert\_start \le x<expert\_end \ \}|83 availableIdxNum = |\{x\in expertIdx| expert\_start \le x<expert\_end \ \}|
84 $$84 $$
85 85 
86-- 等价计算逻辑86+- 等价计算逻辑
87+ 
87 ```python88 ```python
88 import numpy as np89 import numpy as np
89 import random90 import random
@@ -593,76 +594,75 @@
593 demo_drop_pad_mode()594 demo_drop_pad_mode()
594 ```595 ```
595 596 
596- 
597## 函数原型<a name="zh-cn_topic_0000002271534921_section14509346133618"></a>597## 函数原型<a name="zh-cn_topic_0000002271534921_section14509346133618"></a>
598 598 
599-```599+```python
600torch_npu.npu_moe_init_routing_v2(x, expert_idx, *, scale=None, offset=None, active_num=-1, expert_capacity=-1, expert_num=-1, drop_pad_mode=0, expert_tokens_num_type=0, expert_tokens_num_flag=False, quant_mode=-1, active_expert_range=[], row_idx_type=0) -> (Tensor, Tensor, Tensor, Tensor)600torch_npu.npu_moe_init_routing_v2(x, expert_idx, *, scale=None, offset=None, active_num=-1, expert_capacity=-1, expert_num=-1, drop_pad_mode=0, expert_tokens_num_type=0, expert_tokens_num_flag=False, quant_mode=-1, active_expert_range=[], row_idx_type=0) -> (Tensor, Tensor, Tensor, Tensor)
601```601```
602 602 
603## 参数说明<a name="zh-cn_topic_0000002271534921_section2050919466367"></a>603## 参数说明<a name="zh-cn_topic_0000002271534921_section2050919466367"></a>
604 604 
605-- **x** (`Tensor`):必选参数,表示MoE的输入即token特征输入,要求为2维张量,shape为(NUM_ROWS, H)。数据类型支持`float16`、`bfloat16`、`float32`、`int8`,数据格式要求为$ND$。605+- **x** (`Tensor`):必选参数,表示MoE的输入即token特征输入,要求为2维张量,shape为(NUM_ROWS, H)。数据类型支持`float16`、`bfloat16`、`float32`、`int8`,数据格式要求为$ND$。
606-- **expert_idx** (`Tensor`):必选参数,表示[torch_npu.npu_moe_gating_top_k_softmax](torch_npu-npu_moe_gating_top_k_softmax.md)输出每一行特征对应的K个处理专家,要求是2维张量,shape为(NUM_ROWS, K),且专家id不能超过专家数。数据类型支持`int32`,数据格式要求为$ND$。606+- **expert_idx** (`Tensor`):必选参数,表示[torch_npu.npu_moe_gating_top_k_softmax](torch_npu-npu_moe_gating_top_k_softmax.md)输出每一行特征对应的K个处理专家,要求是2维张量,shape为(NUM_ROWS, K),且专家id不能超过专家数。数据类型支持`int32`,数据格式要求为$ND$。
607- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。607- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
608-- **scale** (`Tensor`):可选参数,默认为None,用于计算量化结果的参数。数据类型支持`float32`,数据格式要求为$ND$。如果不输入表示计算时不使用`scale`,且输出`expanded_scale`中的值无意义。608+- **scale** (`Tensor`):可选参数,默认为None,用于计算量化结果的参数。数据类型支持`float32`,数据格式要求为$ND$。如果不输入表示计算时不使用`scale`,且输出`expanded_scale`中的值无意义。
609- - 非量化场景下,如果输入则要求为1维张量,shape为(NUM_ROWS,)。609+ - 非量化场景下,如果输入则要求为1维张量,shape为(NUM_ROWS,)。
610- - 静态量化场景必须输入,输入要求为1D的Tensor,shape为(1,)610+ - 静态量化场景必须输入,输入要求为1D的Tensor,shape为(1,)
611- - 动态量化场景下,如果输入则要求为2维张量,shape为(expert_end-expert_start, H)或(1, H)。611+ - 动态量化场景下,如果输入则要求为2维张量,shape为(expert_end-expert_start, H)或(1, H)。
612 612 
613-- **offset** (`Tensor`):可选参数,默认为None,用于计算量化结果的偏移值。数据类型支持`float32`,数据格式要求为$ND$。613+- **offset** (`Tensor`):可选参数,默认为None,用于计算量化结果的偏移值。数据类型支持`float32`,数据格式要求为$ND$。
614- - 在非量化场景下不输入。614+ - 在非量化场景下不输入。
615- - 静态量化场景必须输入,输入要求为1维张量,shape为(1,)615+ - 静态量化场景必须输入,输入要求为1维张量,shape为(1,)
616- - 动态量化场景下不输入。616+ - 动态量化场景下不输入。
617 617 
618-- **active_num** (`int`):可选参数,默认值为-1,表示总的最大处理row数,输出`expanded_x`只有这么多行是有效的,入参校验需大于等于0,0表示Dropless场景,大于0时表示Active场景,约束所有专家共同处理tokens总量。618+- **active_num** (`int`):可选参数,默认值为-1,表示总的最大处理row数,输出`expanded_x`只有这么多行是有效的,入参校验需大于等于0,0表示Dropless场景,大于0时表示Active场景,约束所有专家共同处理tokens总量。
619-- **expert_capacity** (`int`):可选参数,默认值为-1,表示每个专家能够处理的tokens数,入参校验大于0小于NUM_ROWS。619+- **expert_capacity** (`int`):可选参数,默认值为-1,表示每个专家能够处理的tokens数,入参校验大于0小于NUM_ROWS。
620-- **expert_num** (`int`):可选参数,默认值为-1,表示专家数。`expert_tokens_num_type`为key_value模式时,取值范围为[0, 5120];其他模式取值范围为[0, 10240]。620+- **expert_num** (`int`):可选参数,默认值为-1,表示专家数。`expert_tokens_num_type`为key_value模式时,取值范围为[0, 5120];其他模式取值范围为[0, 10240]。
621-- **drop_pad_mode** (`int`):可选参数,默认值为0,表示是否为drop_pad场景。0表示dropless场景,该场景下不校验`expert_capacity`。1表示drop_pad场景。621+- **drop_pad_mode** (`int`):可选参数,默认值为0,表示是否为drop_pad场景。0表示dropless场景,该场景下不校验`expert_capacity`。1表示drop_pad场景。
622-- **expert_tokens_num_type** (`int`):可选参数,默认值为0,表示直方图的不同模式。取值为0、1和2。0表示cumsum模式;1表示count模式,即输出的值为各个专家处理的token数量的累计值;2表示key_value模式,即输出的值为专家和对应专家处理token数量的累计值。622+- **expert_tokens_num_type** (`int`):可选参数,默认值为0,表示直方图的不同模式。取值为0、1和2。0表示cumsum模式;1表示count模式,即输出的值为各个专家处理的token数量的累计值;2表示key_value模式,即输出的值为专家和对应专家处理token数量的累计值。
623-- **expert_tokens_num_flag** (`bool`):可选参数,默认值为False,取值为False和True,表示是否输出`expert_token_cumsum_or_count`。623+- **expert_tokens_num_flag** (`bool`):可选参数,默认值为False,取值为False和True,表示是否输出`expert_token_cumsum_or_count`。
624-- **quant_mode** (`int`):可选参数,默认值为-1,表示量化模式,支持取值为0、1、-1。0表示静态量化,-1表示不量化场景;1表示动态量化场景。624+- **quant_mode** (`int`):可选参数,默认值为-1,表示量化模式,支持取值为0、1、-1。0表示静态量化,-1表示不量化场景;1表示动态量化场景。
625-- **active_expert_range** (`List[int]`):可选参数,默认为空, 表示活跃expert的范围。数组内值的范围为[expert_start, expert_end],左闭右开,表示活跃的expert范围在expert_start到expert_end之间。要求值大于等于0,并且expert_end不大于`expert_num`。drop_pad场景下,expert_start等于0, expert_end等于`expert_num`。传入默认值时,视为活跃的expert范围在0到`expert_num`之间。625+- **active_expert_range** (`List[int]`):可选参数,默认为空, 表示活跃expert的范围。数组内值的范围为[expert_start, expert_end],左闭右开,表示活跃的expert范围在expert_start到expert_end之间。要求值大于等于0,并且expert_end不大于`expert_num`。drop_pad场景下,expert_start等于0, expert_end等于`expert_num`。传入默认值时,视为活跃的expert范围在0到`expert_num`之间。
626-- **row_idx_type** (`int`):可选参数,默认为0,表示输出`expanded_row_idx`使用的索引类型,支持取值0和1。0表示gather类型的索引;1表示scatter类型的索引。626+- **row_idx_type** (`int`):可选参数,默认为0,表示输出`expanded_row_idx`使用的索引类型,支持取值0和1。0表示gather类型的索引;1表示scatter类型的索引。
627 627 
628## 返回值说明<a name="zh-cn_topic_0000002271534921_section18510124618368"></a>628## 返回值说明<a name="zh-cn_topic_0000002271534921_section18510124618368"></a>
629 629 
630-- **expanded_x** (`Tensor`):根据`expert_idx`进行扩展过的特征,Dropless场景shape为[NUM_ROWS * K, H]。Active场景shape为[min(activeNum, NUM_ROWS * K), H]。Drop/Pad场景下要求是一个3D的Tensor,shape为[expertNum, expertCapacity, H]。非量化场景下数据类型同`x`;量化场景下数据类型为`int8`。数据格式要求为$ND$。量化场景下,当`x`的数据类型为`int8`时,输出值无意义。630+- **expanded_x** (`Tensor`):根据`expert_idx`进行扩展过的特征,Dropless场景shape为[NUM_ROWS \* K, H]。Active场景shape为[min(activeNum, NUM_ROWS * K), H]。Drop/Pad场景下要求是一个3D的Tensor,shape为[expertNum, expertCapacity, H]。非量化场景下数据类型同`x`;量化场景下数据类型为`int8`。数据格式要求为$ND$。量化场景下,当`x`的数据类型为`int8`时,输出值无意义。
631-- **expanded_row_idx** (`Tensor`):`expanded_x`和`x`的映射关系,要求是1维张量,shape为(NUM_ROWS\*K, ),数据类型支持`int32`,数据格式要求为$ND$。当`row_idx_type`为1时, 前available_idx_num个元素为有效数据,无效数据未初始化;当`row_idx_type`为0时,无效数据由-1填充。631+- **expanded_row_idx** (`Tensor`):`expanded_x`和`x`的映射关系,要求是1维张量,shape为(NUM_ROWS \* K, ),数据类型支持`int32`,数据格式要求为$ND$。当`row_idx_type`为1时, 前available_idx_num个元素为有效数据,无效数据未初始化;当`row_idx_type`为0时,无效数据由-1填充。
632-- **expert_token_cumsum_or_count** (`Tensor`):表示输出每个专家处理的token数量的统计结果或累加值。632+- **expert_token_cumsum_or_count** (`Tensor`):表示输出每个专家处理的token数量的统计结果或累加值。
633- - 在`expert_tokens_num_type`为0时,表示`active_expert_range`范围内expert在排序后处理token总数的前缀和。633+ - 在`expert_tokens_num_type`为0时,表示`active_expert_range`范围内expert在排序后处理token总数的前缀和。
634- - 在`expert_tokens_num_type`为1的场景下,要求是1维张量,表示`active_expert_range`范围内expert对应的处理token的总数,shape为(expert_end-expert_start, )。shape为(expert_end-expert_start, );634+ - 在`expert_tokens_num_type`为1的场景下,要求是1维张量,表示`active_expert_range`范围内expert对应的处理token的总数,shape为(expert_end-expert_start, )。shape为(expert_end-expert_start, );
635- - 在`expert_tokens_num_type`为2的场景下,要求是2维张量,shape为(expert_num, 2),表示`active_expert_range`范围内token总数为非0的expert,以及对应expert处理token的总数;635+ - 在`expert_tokens_num_type`为2的场景下,要求是2维张量,shape为(expert_num, 2),表示`active_expert_range`范围内token总数为非0的expert,以及对应expert处理token的总数;
636 636 
637 expert_idx在active_expert_range范围且剔除对应expert处理token为0的元素对为有效元素对,存放于Tensor头部并保持原序。数据类型支持`int64`,数据格式要求为$ND$。637 expert_idx在active_expert_range范围且剔除对应expert处理token为0的元素对为有效元素对,存放于Tensor头部并保持原序。数据类型支持`int64`,数据格式要求为$ND$。
638-- **expanded_scale** (`Tensor`):数据类型支持`float32`,数据格式要求为$ND$。输出shape为`expert_idx`的shape去掉最后一维之后所有维度的乘积。令available_idx_num为`active_expert_range`范围的元素的个数。638+- **expanded_scale** (`Tensor`):数据类型支持`float32`,数据格式要求为$ND$。输出shape为`expert_idx`的shape去掉最后一维之后所有维度的乘积。令available_idx_num为`active_expert_range`范围的元素的个数。
639- - 非量化场景下,当`scale`输入时,前`available_idx_num`个元素为有效数据。639+ - 非量化场景下,当`scale`输入时,前`available_idx_num`个元素为有效数据。
640- - 动态量化场景下,输出量化计算过程中`scale`的中间值,前`available_idx_num`个元素为有效数据。640+ - 动态量化场景下,输出量化计算过程中`scale`的中间值,前`available_idx_num`个元素为有效数据。
641- - 静态量化场景下不输出。641+ - 静态量化场景下不输出。
642 642 
643## 约束说明<a name="zh-cn_topic_0000002271534921_section75102046193618"></a>643## 约束说明<a name="zh-cn_topic_0000002271534921_section75102046193618"></a>
644 644 
645-- 该接口支持推理场景下使用。645+- 该接口支持推理场景下使用。
646-- 该接口支持图模式。646+- 该接口支持图模式。
647-- 进入低时延性能模板需要同时满足以下条件:647+- 进入低时延性能模板需要同时满足以下条件:
648- - `x`、`expert_idx`、`scale`输入Shape要求分别为:(1, 7168)、(1, 8)、(256, 7168)648+ - `x`、`expert_idx`、`scale`输入Shape要求分别为:(1, 7168)、(1, 8)、(256, 7168)
649- - `x`数据类型要求:`bfloat16`649+ - `x`数据类型要求:`bfloat16`
650- - 属性要求:active_expert_range=[0, 256]、 quant_mode=1、expert_tokens_num_type=2、expert_num=256650+ - 属性要求:active_expert_range=[0, 256]、 quant_mode=1、expert_tokens_num_type=2、expert_num=256
651 651 
652-- 进入大batch性能模板需要同时满足以下条件:652+- 进入大batch性能模板需要同时满足以下条件:
653- - NUM_ROWS范围为[384, 8192]653+ - NUM_ROWS范围为[384, 8192]
654- - K=8654+ - K=8
655- - expert_num=256655+ - expert_num=256
656- - expert_end-expert_start<=32656+ - expert_end-expert_start<=32
657- - quant_mode=-1657+ - quant_mode=-1
658- - row_idx_type=1658+ - row_idx_type=1
659- - expert_tokens_num_type=1659+ - expert_tokens_num_type=1
660 660 
661-- 在算子输入shape较小的场景,操作间的多核同步时间占比较高,成为性能瓶颈。因此,针对这种特化场景,添加全载性能模板。该模板中,搬入、排序、计算都在同一个kernel内完成。需要满足 drop_pad_mode=0 的条件。661+- 在算子输入shape较小的场景,操作间的多核同步时间占比较高,成为性能瓶颈。因此,针对这种特化场景,添加全载性能模板。该模板中,搬入、排序、计算都在同一个kernel内完成。需要满足 drop_pad_mode=0 的条件。
662 662 
663## 调用示例<a name="zh-cn_topic_0000002271534921_section12510194643618"></a>663## 调用示例<a name="zh-cn_topic_0000002271534921_section12510194643618"></a>
664 664 
665-- 单算子模式调用665+- 单算子模式调用
666 666 
667 ```python667 ```python
668 import torch668 import torch
@@ -693,7 +693,7 @@ torch_npu.npu_moe_init_routing_v2(x, expert_idx, *, scale=None, offset=None, act
693 active_expert_range=active_expert_range, quant_mode=quant_mode, row_idx_type=row_idx_type)693 active_expert_range=active_expert_range, quant_mode=quant_mode, row_idx_type=row_idx_type)
694 ```694 ```
695 695 
696-- 图模式调用696+- 图模式调用
697 697 
698 ```python698 ```python
699 import torch699 import torch
@@ -748,5 +748,3 @@ torch_npu.npu_moe_init_routing_v2(x, expert_idx, *, scale=None, offset=None, act
748 if __name__ == '__main__':748 if __name__ == '__main__':
749 main()749 main()
750 ```750 ```
751- 
752- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_re_routing.md+25-25
@@ -9,19 +9,19 @@
9 9 
10## 功能说明10## 功能说明
11 11 
12-- API功能:MoE网络中,进行AlltoAll操作从其他卡上拿到需要算的token后,将token按照专家顺序重新排列。12+- API功能:MoE网络中,进行AlltoAll操作从其他卡上拿到需要算的token后,将token按照专家顺序重新排列。
13-- 计算公式:13+- 计算公式:
14 14 
15 ![](../../figures/zh-cn_formulaimage_0000002277237821.png)15 ![](../../figures/zh-cn_formulaimage_0000002277237821.png)
16 16 
17- - SrcOffset指当前需要移动的token源偏移,根据输入`expert_token_num_per_rank`的值进行计算。17+ - SrcOffset指当前需要移动的token源偏移,根据输入`expert_token_num_per_rank`的值进行计算。
18- - DstOffset指当前需要移动的token目的偏移。18+ - DstOffset指当前需要移动的token目的偏移。
19- - cur\_rank是`expert_token_num_per_rank`的纵轴索引,表示该token原本在的卡。19+ - cur\_rank是`expert_token_num_per_rank`的纵轴索引,表示该token原本在的卡。
20- - cur\_expert是`expert_token_num_per_rank`的横轴索引,表示该token由卡上专家cur\_expert计算。20+ - cur\_expert是`expert_token_num_per_rank`的横轴索引,表示该token由卡上专家cur\_expert计算。
21 21 
22## 函数原型22## 函数原型
23 23 
24-```24+```python
25torch_npu.npu_moe_re_routing(tokens, expert_token_num_per_rank, *, per_token_scales=None, expert_token_num_type=1, idx_type=0) -> (Tensor, Tensor, Tensor, Tensor)25torch_npu.npu_moe_re_routing(tokens, expert_token_num_per_rank, *, per_token_scales=None, expert_token_num_type=1, idx_type=0) -> (Tensor, Tensor, Tensor, Tensor)
26```26```
27 27 
@@ -29,33 +29,34 @@ torch_npu.npu_moe_re_routing(tokens, expert_token_num_per_rank, *, per_token_sca
29 29 
30> [!NOTE] 30> [!NOTE]
31> Tensor中shape使用的变量说明:31> Tensor中shape使用的变量说明:
32-> - A:表示token个数,取值要求Sum\(expert\_token\_num\_per\_rank\)=A。32+>
33-> - H:表示token长度,取值要求0<H<1638433+> - A:表示token个数,取值要求Sum\(expert\_token\_num\_per\_rank\)=A
34-> - N:表示卡数,取值无限制34+> - H:表示token长度,取值要求0<H<16384
35-> - E:表示卡上的专家数,取值无限制。35+> - N:表示卡数,取值无限制。
36+> - E:表示卡上的专家数,取值无限制。
36 37 
37-- **tokens** (`Tensor`):必选参数,表示待重新排布的token。要求为2维,shape为\[A, H\],数据类型支持`float16`、`bfloat16`、`int8`,数据格式为$ND$。38+- **tokens** (`Tensor`):必选参数,表示待重新排布的token。要求为2维,shape为\[A, H\],数据类型支持`float16`、`bfloat16`、`int8`,数据格式为$ND$。
38-- **expert\_token\_num\_per\_rank** (`Tensor`):必选参数,二维矩阵,矩阵中元素[i, j]表示当前卡上从卡i获取到的专家j处理的token数。要求为2维,shape为\[N, E\],数据类型支持`int32`、`int64`,数据格式为$ND$。取值必须大于0。39+- **expert\_token\_num\_per\_rank** (`Tensor`):必选参数,二维矩阵,矩阵中元素[i, j]表示当前卡上从卡i获取到的专家j处理的token数。要求为2维,shape为\[N, E\],数据类型支持`int32`、`int64`,数据格式为$ND$。取值必须大于0。
39- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。40- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
40-- **per\_token\_scales:** (`Tensor`):可选参数,表示每个token对应的scale,需要随token同样进行重新排布。要求为1维,shape为\[A\],数据类型支持`float32`,数据格式为$ND$。41+- **per\_token\_scales:** (`Tensor`):可选参数,表示每个token对应的scale,需要随token同样进行重新排布。要求为1维,shape为\[A\],数据类型支持`float32`,数据格式为$ND$。
41-- **expert\_token\_num\_type** (`int`):可选参数,表示输出`expert_token_num`的模式。0为cumsum模式,1为count模式,默认值为1。当前只支持为1。42+- **expert\_token\_num\_type** (`int`):可选参数,表示输出`expert_token_num`的模式。0为cumsum模式,1为count模式,默认值为1。当前只支持为1。
42-- **idx\_type** (`int`):可选参数,表示输出`permute_token_idx`的索引类型。0为gather索引,1为scatter索引,默认值为0。当前只支持为0。43+- **idx\_type** (`int`):可选参数,表示输出`permute_token_idx`的索引类型。0为gather索引,1为scatter索引,默认值为0。当前只支持为0。
43 44 
44## 返回值说明45## 返回值说明
45 46 
46-- **permute\_tokens** (`Tensor`):表示重新排布后的token。要求为2维,shape为\[A, H\],数据类型支持`float16`、`bfloat16`、`int8`,数据格式为$ND$。47+- **permute\_tokens** (`Tensor`):表示重新排布后的token。要求为2维,shape为\[A, H\],数据类型支持`float16`、`bfloat16`、`int8`,数据格式为$ND$。
47-- **permute\_per\_token\_scales** (`Tensor`):表示重新排布后的`per_token_scales`,输入不携带`per_token_scales`的情况下,该输出无效。要求为1维,shape为\[A\],数据类型支持`float32`,数据格式为$ND$。48+- **permute\_per\_token\_scales** (`Tensor`):表示重新排布后的`per_token_scales`,输入不携带`per_token_scales`的情况下,该输出无效。要求为1维,shape为\[A\],数据类型支持`float32`,数据格式为$ND$。
48-- **permute\_token\_idx** (`Tensor`):表示每个token在原排布方式的索引。要求为1维,shape为\[A\],数据类型支持`int32`,数据格式为$ND$。49+- **permute\_token\_idx** (`Tensor`):表示每个token在原排布方式的索引。要求为1维,shape为\[A\],数据类型支持`int32`,数据格式为$ND$。
49-- **expert\_token\_num** (`Tensor`):表示每个专家处理的token数。要求为1维,shape为\[E\],数据类型支持`int32`、`int64`,数据格式为$ND$。50+- **expert\_token\_num** (`Tensor`):表示每个专家处理的token数。要求为1维,shape为\[E\],数据类型支持`int32`、`int64`,数据格式为$ND$。
50 51 
51## 约束说明52## 约束说明
52 53 
53-- 该接口支持推理场景下使用。54+- 该接口支持推理场景下使用。
54-- 该接口支持图模式。55+- 该接口支持图模式。
55 56 
56## 调用示例57## 调用示例
57 58 
58-- 单算子模式调用59+- 单算子模式调用
59 60 
60 ```python61 ```python
61 import torch62 import torch
@@ -96,7 +97,7 @@ torch_npu.npu_moe_re_routing(tokens, expert_token_num_per_rank, *, per_token_sca
96 permute_tokens_npu, permute_per_token_scales_npu, permute_token_idx_npu, expert_token_num_npu = torch_npu.npu_moe_re_routing(tokens_npu, expert_token_num_per_rank_npu, per_token_scales=per_token_scales_npu, expert_token_num_type=expert_token_num_type, idx_type=idx_type)97 permute_tokens_npu, permute_per_token_scales_npu, permute_token_idx_npu, expert_token_num_npu = torch_npu.npu_moe_re_routing(tokens_npu, expert_token_num_per_rank_npu, per_token_scales=per_token_scales_npu, expert_token_num_type=expert_token_num_type, idx_type=idx_type)
97 ```98 ```
98 99 
99-- 图模式调用100+- 图模式调用
100 101 
101 ```python102 ```python
102 import torch103 import torch
@@ -151,4 +152,3 @@ torch_npu.npu_moe_re_routing(tokens, expert_token_num_per_rank, *, per_token_sca
151 model = torch.compile(model, backend=npu_backend, dynamic=False)152 model = torch.compile(model, backend=npu_backend, dynamic=False)
152 permute_tokens_npu, permute_per_token_scales_npu, permute_token_idx_npu, expert_token_num_npu = model(tokens_npu, expert_token_num_per_rank_npu, per_token_scales_npu, expert_token_num_type, idx_type)153 permute_tokens_npu, permute_per_token_scales_npu, permute_token_idx_npu, expert_token_num_npu = model(tokens_npu, expert_token_num_per_rank_npu, per_token_scales_npu, expert_token_num_type, idx_type)
153 ```154 ```
154- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_moe_update_expert.md+30-24
@@ -6,12 +6,14 @@
6| :----------------------------------------------------------- | :------: |6| :----------------------------------------------------------- | :------: |
7| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |7| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |
8 8 
9- 
10## 功能说明9## 功能说明
10+ 
11本API支持负载均衡和专家剪枝功能。经过映射后的专家表和mask可传入Moe层进行数据分发和处理。11本API支持负载均衡和专家剪枝功能。经过映射后的专家表和mask可传入Moe层进行数据分发和处理。
12+ 
12- 负载均衡:完成冗余专家部署场景下每个token的topK个专家逻辑卡号到物理卡号的映射。计算方法如下所示:13- 负载均衡:完成冗余专家部署场景下每个token的topK个专家逻辑卡号到物理卡号的映射。计算方法如下所示:
13 14 
14 负载均衡对于`expert_ids`中的第i个值,即第i个token:15 负载均衡对于`expert_ids`中的第i个值,即第i个token:
16+ 
15 ```python17 ```python
16 new_expert_id = eplb_table[table_offset + 1]18 new_expert_id = eplb_table[table_offset + 1]
17 expert_id = expert_ids[i]19 expert_id = expert_ids[i]
@@ -27,9 +29,11 @@
27 place_idx = i % place_num29 place_idx = i % place_num
28 new_expert_id = eplb_table[table_offset + place_idx]30 new_expert_id = eplb_table[table_offset + place_idx]
29 ```31 ```
32+ 
30- 专家剪枝:支持根据阈值对token发送的topK个专家进行剪枝。计算方法如下所示:33- 专家剪枝:支持根据阈值对token发送的topK个专家进行剪枝。计算方法如下所示:
31 34
32 将shape为$(BS,)$的`active_mask`进行broadcast成为shape为$(BS,K)$的active_mask_tensor,其中BS对应为False的专家会直接被剪枝。对于active_mask_tensor为True的`expert_scales`的元素,满足条件也将被剪枝。35 将shape为$(BS,)$的`active_mask`进行broadcast成为shape为$(BS,K)$的active_mask_tensor,其中BS对应为False的专家会直接被剪枝。对于active_mask_tensor为True的`expert_scales`的元素,满足条件也将被剪枝。
36+ 
33 ```python37 ```python
34 active_mask_tensor = broadcast(active_mask, (BS, K))38 active_mask_tensor = broadcast(active_mask, (BS, K))
35 for i in range(BS):39 for i in range(BS):
@@ -39,31 +43,32 @@
39 43 
40## 函数原型44## 函数原型
41 45 
42-```46+```python
43torch_npu.npu_moe_update_expert(expert_ids, eplb_table, *, expert_scales=None, pruning_threshold=None, active_mask=None, local_rank_id=-1, world_size=-1, balance_mode=0) -> (Tensor, Tensor)47torch_npu.npu_moe_update_expert(expert_ids, eplb_table, *, expert_scales=None, pruning_threshold=None, active_mask=None, local_rank_id=-1, world_size=-1, balance_mode=0) -> (Tensor, Tensor)
44```48```
45 49 
46## 参数说明50## 参数说明
47 51 
48-- **expert_ids**(`Tensor`):必选参数,表示每个token的topK个专家索引,shape为$(BS, K)$。数据类型支持`int32`、`int64`,数据格式要求为$ND$,支持非连续的Tensor。52+- **expert_ids**(`Tensor`):必选参数,表示每个token的topK个专家索引,shape为$(BS, K)$。数据类型支持`int32`、`int64`,数据格式要求为$ND$,支持非连续的Tensor。
49-- **eplb_table**(`Tensor`):必选参数,表示逻辑专家到物理专家的映射表,外部调用者需保证输入Tensor的值正确:每行第一列为行号对应逻辑专家部署的实例数count,值需大于等于1,每行\[1, count\]列为对应实例的卡号,取值范围\[0, `moe_expert_num`\),shape为$(moe\_expert\_num, F)$。数据类型支持`int32`,数据格式要求为$ND$,支持非连续的Tensor。其中F表示输入映射表的列数,取值范围\[2, `world_size`+1\],第一列为各行号对应Moe专家部署的实例个数(值>0),后F-1列为该Moe专家部署的物理卡号。53+- **eplb_table**(`Tensor`):必选参数,表示逻辑专家到物理专家的映射表,外部调用者需保证输入Tensor的值正确:每行第一列为行号对应逻辑专家部署的实例数count,值需大于等于1,每行\[1, count\]列为对应实例的卡号,取值范围\[0, `moe_expert_num`\),shape为$(moe\_expert\_num, F)$。数据类型支持`int32`,数据格式要求为$ND$,支持非连续的Tensor。其中F表示输入映射表的列数,取值范围\[2, `world_size`+1\],第一列为各行号对应Moe专家部署的实例个数(值>0),后F-1列为该Moe专家部署的物理卡号。
50-- **expert_scales**(`Tensor`):可选参数,每个token的topK个专家的scale权重,用户需保证scale在token内部按照降序排列,可选择传入有效数据或空指针,该参数传入有效数据时,`pruning_threshold`也需要传入有效数据。shape为$(BS, K)$。数据类型支持`fp16`、`bf16`、`float`,数据格式要求为$ND$,支持非连续的Tensor。54+- **expert_scales**(`Tensor`):可选参数,每个token的topK个专家的scale权重,用户需保证scale在token内部按照降序排列,可选择传入有效数据或空指针,该参数传入有效数据时,`pruning_threshold`也需要传入有效数据。shape为$(BS, K)$。数据类型支持`fp16`、`bf16`、`float`,数据格式要求为$ND$,支持非连续的Tensor。
51-- **pruning_threshold**(`Tensor`):可选参数,专家scale权重的最小阈值,当某个token对应的某个topK专家scale小于阈值时,该token将对该专家进行剪枝,即token不发送至该专家处理,可选择传入有效数据或空指针,该参数传入有效数据时,`expert_scales`也需要传入有效数据。shape为$(K,)$或$(1, K)$。数据类型支持`float`,数据格式要求为$ND$,支持非连续的Tensor。55+- **pruning_threshold**(`Tensor`):可选参数,专家scale权重的最小阈值,当某个token对应的某个topK专家scale小于阈值时,该token将对该专家进行剪枝,即token不发送至该专家处理,可选择传入有效数据或空指针,该参数传入有效数据时,`expert_scales`也需要传入有效数据。shape为$(K,)$或$(1, K)$。数据类型支持`float`,数据格式要求为$ND$,支持非连续的Tensor。
52-- **active_mask**(`Tensor`):可选参数,表示token是否参与通信,可选择传入有效数据或空指针。传入有效数据时,`expert_scales`、`pruning_threshold`也必须传入有效数据,参数为true表示对应的token参与通信,true必须排到false之前,例:\{true, false, true\}为非法输入;传入空指针时表示所有token都会参与通信。shape为$(BS,)$。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。56+- **active_mask**(`Tensor`):可选参数,表示token是否参与通信,可选择传入有效数据或空指针。传入有效数据时,`expert_scales`、`pruning_threshold`也必须传入有效数据,参数为true表示对应的token参与通信,true必须排到false之前,例:\{true, false, true\}为非法输入;传入空指针时表示所有token都会参与通信。shape为$(BS,)$。数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。
53 57 
54-- **local_rank_id**(`int`):本卡ID,数据类型支持`int64`,当`balance_mode`设置0时,本属性取值范围为\[0, `world_size`\)。58+- **local_rank_id**(`int`):本卡ID,数据类型支持`int64`,当`balance_mode`设置0时,本属性取值范围为\[0, `world_size`\)。
55-- **world_size**(`int`):通信域size,数据类型支持`int64`,当`balance_mode`设置0时,本属性取值范围为\[2, 768\]59+- **world_size**(`int`):通信域size,数据类型支持`int64`,当`balance_mode`设置0时,本属性取值范围为\[2, 768\]
56-- **balance_mode**(`int`):均衡规则,数据类型支持`int64`,取值支持0和1,0表示用`local_rank_id`进行负载均衡,1表示使用`token_id`进行负载均衡。当本属性取值为0时,`local_rank_id`和`world_size`必须传入有效值。60+- **balance_mode**(`int`):均衡规则,数据类型支持`int64`,取值支持0和1,0表示用`local_rank_id`进行负载均衡,1表示使用`token_id`进行负载均衡。当本属性取值为0时,`local_rank_id`和`world_size`必须传入有效值。
57 61 
58## 返回值说明62## 返回值说明
59 63 
60-- **balanced_expert_ids**(`Tensor`):映射后每个token的topK个专家所在物理卡的卡号,shape为(BS, K),数据类型、数据格式与`expert_ids`保持一致。64+- **balanced_expert_ids**(`Tensor`):映射后每个token的topK个专家所在物理卡的卡号,shape为(BS, K),数据类型、数据格式与`expert_ids`保持一致。
61-- **balanced_active_mask**(`Tensor`):剪枝后的`active_mask`,当`expert_scales`、`pruning_threshold`传入有效数据时该输出有效。shape为\(BS, K\),数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。65+- **balanced_active_mask**(`Tensor`):剪枝后的`active_mask`,当`expert_scales`、`pruning_threshold`传入有效数据时该输出有效。shape为\(BS, K\),数据类型支持`bool`,数据格式要求为$ND$,支持非连续的Tensor。
62 66 
63## 约束说明67## 约束说明
64 68 
65-- 该接口必须与`torch_npu.npu_moe_distribute_dispatch`或`torch_npu.npu_moe_distribute_dispatch_v2`接口配合使用。69+- 该接口必须与`torch_npu.npu_moe_distribute_dispatch`或`torch_npu.npu_moe_distribute_dispatch_v2`接口配合使用。
66-- 调用接口过程中使用的`world_size` 、`moe_expert_num` 参数取值所有卡须保持一致,网络中不同层中也需保持一致,本接口中参数和`torch_npu.npu_moe_distribute_dispatch`或`torch_npu.npu_moe_distribute_dispatch_v2`有如下对应关系:70+- 调用接口过程中使用的`world_size` 、`moe_expert_num` 参数取值所有卡须保持一致,网络中不同层中也需保持一致,本接口中参数和`torch_npu.npu_moe_distribute_dispatch`或`torch_npu.npu_moe_distribute_dispatch_v2`有如下对应关系:
71+ 
67 |`torch_npu.npu_moe_update_expert` |`torch_npu.npu_moe_distribute_dispatch`/`torch_npu.npu_moe_distribute_dispatch_v2`|72 |`torch_npu.npu_moe_update_expert` |`torch_npu.npu_moe_distribute_dispatch`/`torch_npu.npu_moe_distribute_dispatch_v2`|
68 |---------------------------|-----------------------------------------------------------|73 |---------------------------|-----------------------------------------------------------|
69 |`local_rank_id` |`ep_rank_id` |74 |`local_rank_id` |`ep_rank_id` |
@@ -71,18 +76,19 @@ torch_npu.npu_moe_update_expert(expert_ids, eplb_table, *, expert_scales=None, p
71 |`eplb_table`第一列的count之和 |`moe_expert_num` |76 |`eplb_table`第一列的count之和 |`moe_expert_num` |
72 |`BS` |`BS` |77 |`BS` |`BS` |
73 |`K` |`K` |78 |`K` |`K` |
74-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。79+ 
75-- 参数说明里shape格式说明:80+- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:该场景下单卡包含双DIE(简称为“晶粒”或“裸片”),因此参数说明里的“本卡”均表示单DIE。
76- - `BS`:表示batch sequence size,即本卡最终输出的token量,<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>取值范围为0<BS≤512。81+- 说明里shape格式说明
77- - `K`:表示选取topK个专家,取值范围为0< K 16同时满足0 < K ≤ moe_expert_num82+ - `BS`:表示batch sequence size即本卡最终输出的token数量,<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:取值范围为0<BS512
78- - `moe_expert_num`:表示Moe专家,取值范围\(0, 1024\]83+ - `K`:表示选取topK个专家,取值范围0< K ≤16同时满足0 < K ≤ moe_expert_num
79- - `F`:表示输入映射表`eplb_table`的列数,取值范围\[2, world_size + 1\]。84+ - `moe_expert_num`:表示Moe专家数,取值范围\(0, 1024\]。
80- - 每个专家部署副本个数值(即eplb_table第一列count)最小1,最大为`world_size`85+ - `F`:表示输入映射表`eplb_table`列数取值范围\[2, world_size + 1\]
81- - 所有专家部署副本个数(即eplb_table第一列count之和于等于1024且整除`world_size`。86+ - 每个专家部署副本个数(即eplb_table第一列count),最为1最大为`world_size`。
87+ - 所有专家部署的副本个数和(即eplb_table第一列count之和)需小于等于1024,且整除`world_size`。
82 88 
83## 调用示例89## 调用示例
84 90 
85-- 单算子模式调用91+- 单算子模式调用
86 92 
87 ```python93 ```python
88 import os94 import os
@@ -272,7 +278,7 @@ torch_npu.npu_moe_update_expert(expert_ids, eplb_table, *, expert_scales=None, p
272 print("run npu success.")278 print("run npu success.")
273 ```279 ```
274 280 
275-- 图模式调用281+- 图模式调用
276 282 
277 ```python283 ```python
278 # 修改graph_type支持静态图、动态图284 # 修改graph_type支持静态图、动态图
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_prefetch.md+1-6
@@ -1,6 +1,5 @@
1# torch_npu.npu_prefetch1# torch_npu.npu_prefetch
2 2 
3- 
4## 产品支持情况3## 产品支持情况
5 4 
6| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -8,14 +7,13 @@
8|<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |7|<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |
9|<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |8|<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |
10 9 
11-
12## 功能说明10## 功能说明
13 11 
14提供网络`weight`预取功能,用于在计算执行前将指定的权重数据预先加载到L2 Cache中,减少算子访问这些权重时的访存等待时间。例如,在MatMul等算子之前进行预取,算子执行时可直接从低时延的L2 Cache中读取权重,进而提升算子数据访问与计算效率。实际性能收益取决于用户采用的并行方式和配置。12提供网络`weight`预取功能,用于在计算执行前将指定的权重数据预先加载到L2 Cache中,减少算子访问这些权重时的访存等待时间。例如,在MatMul等算子之前进行预取,算子执行时可直接从低时延的L2 Cache中读取权重,进而提升算子数据访问与计算效率。实际性能收益取决于用户采用的并行方式和配置。
15 13 
16## 函数原型14## 函数原型
17 15 
18-```16+```python
19torch_npu.npu_prefetch(input, dependency, max_size, offset=0) -> None17torch_npu.npu_prefetch(input, dependency, max_size, offset=0) -> None
20```18```
21 19 
@@ -34,8 +32,6 @@ torch_npu.npu_prefetch(input, dependency, max_size, offset=0) -> None
34 32 
35该接口支持图模式。33该接口支持图模式。
36 34 
37- 
38- 
39## 调用示例35## 调用示例
40 36 
41- 单算子多流并发调用37- 单算子多流并发调用
@@ -114,4 +110,3 @@ torch_npu.npu_prefetch(input, dependency, max_size, offset=0) -> None
114 [75411.2188]], device='npu:0')110 [75411.2188]], device='npu:0')
115 torch.Size([10000, 1])111 torch.Size([10000, 1])
116 ```112 ```
117- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_prompt_flash_attention.md+2-2
@@ -19,7 +19,7 @@ $$
19 19 
20## 函数原型20## 函数原型
21 21 
22-```22+```python
23torch_npu.npu_prompt_flash_attention(query, key, value, *, pse_shift=None, padding_mask=None, atten_mask=None, actual_seq_lengths=None, deq_scale1=None, quant_scale1=None, deq_scale2=None, quant_scale2=None, quant_offset2=None, num_heads=1, scale_value=1.0, pre_tokens=2147483647, next_tokens=0, input_layout="BSH",num_key_value_heads=0, actual_seq_lengths_kv=None, sparse_mode=0) -> Tensor23torch_npu.npu_prompt_flash_attention(query, key, value, *, pse_shift=None, padding_mask=None, atten_mask=None, actual_seq_lengths=None, deq_scale1=None, quant_scale1=None, deq_scale2=None, quant_scale2=None, quant_offset2=None, num_heads=1, scale_value=1.0, pre_tokens=2147483647, next_tokens=0, input_layout="BSH",num_key_value_heads=0, actual_seq_lengths_kv=None, sparse_mode=0) -> Tensor
24```24```
25 25 
@@ -84,6 +84,7 @@ torch_npu.npu_prompt_flash_attention(query, key, value, *, pse_shift=None, paddi
84 - `sparse_mode`为5、6、7、8时,分别代表`prefix、global、dilated、block_local`,均暂不支持。84 - `sparse_mode`为5、6、7、8时,分别代表`prefix、global、dilated、block_local`,均暂不支持。
85 85 
86## 返回值说明86## 返回值说明
87+ 
87`Tensor`88`Tensor`
88 89 
89公式中的$atten\_out$,表示计算的最终结果。当`input_layout`为$BNSD\_BSND$时,输入`query`的shape是$BNSD$,输出shape为$BSND$,其余情况shape与`query`的shape保持一致。90公式中的$atten\_out$,表示计算的最终结果。当`input_layout`为$BNSD\_BSND$时,输入`query`的shape是$BNSD$,输出shape为$BSND$,其余情况shape与`query`的shape保持一致。
@@ -240,4 +241,3 @@ torch_npu.npu_prompt_flash_attention(query, key, value, *, pse_shift=None, paddi
240 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],241 [ 0.0176, 0.0288, -0.0091, ..., 0.0304, 0.0033, -0.0173]]]],
241 device='npu:0', dtype=torch.float16) torch.Size([1, 8, 164, 128])242 device='npu:0', dtype=torch.float16) torch.Size([1, 8, 164, 128])
242 ```243 ```
243- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_quant_lightning_indexer.md+32-25
@@ -1,6 +1,7 @@
1# torch_npu.npu_quant_lightning_indexer1# torch_npu.npu_quant_lightning_indexer
2 2 
3## 产品支持情况3## 产品支持情况
4+ 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
5| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
6|<term>Atlas A3 推理系列产品</term> | √ |7|<term>Atlas A3 推理系列产品</term> | √ |
@@ -8,9 +9,9 @@
8 9 
9## 功能说明10## 功能说明
10 11 
11-- API功能:QuantLightningIndexer是推理场景下,SparseFlashAttention(SFA)前处理的计算,选出关键的稀疏token,并对输入query和key进行量化实现存8算8,获取最大收益。12+- API功能:QuantLightningIndexer是推理场景下,SparseFlashAttention(SFA)前处理的计算,选出关键的稀疏token,并对输入query和key进行量化实现存8算8,获取最大收益。
12 13 
13-- 计算公式:14+- 计算公式:
14 $$out = \text{Top-}k\left\{[1]_{1\times g}@\left[(W@[1]_{1\times S_{k}})\odot\text{ReLU}\left(\left(Scale_Q@Scale_K^T\right)\odot\left(Q_{index}^{INT8}@{\left(K_{index}^{INT8}\right)}^T\right)\right)\right]\right\}$$15 $$out = \text{Top-}k\left\{[1]_{1\times g}@\left[(W@[1]_{1\times S_{k}})\odot\text{ReLU}\left(\left(Scale_Q@Scale_K^T\right)\odot\left(Q_{index}^{INT8}@{\left(K_{index}^{INT8}\right)}^T\right)\right)\right]\right\}$$
15 主要计算过程为:16 主要计算过程为:
16 1. 将某个token对应的输入参数`query`($Q_{index}^{INT8}\in\R^{g\times d}$)乘以给定上下文`key`($K_{index}^{INT8}\in\R^{S_{k}\times d}$),得到相关性。17 1. 将某个token对应的输入参数`query`($Q_{index}^{INT8}\in\R^{g\times d}$)乘以给定上下文`key`($K_{index}^{INT8}\in\R^{S_{k}\times d}$),得到相关性。
@@ -19,62 +20,67 @@
19 20 
20## 函数原型21## 函数原型
21 22 
22-```23+```python
23torch_npu.npu_quant_lightning_indexer(query, key, weights, query_dequant_scale, key_dequant_scale, query_quant_mode, key_quant_mode, *, actual_seq_lengths_query=None, actual_seq_lengths_key=None, block_table=None, layout_query='BSND', layout_key='BSND', sparse_count=2048, sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1) -> Tensor24torch_npu.npu_quant_lightning_indexer(query, key, weights, query_dequant_scale, key_dequant_scale, query_quant_mode, key_quant_mode, *, actual_seq_lengths_query=None, actual_seq_lengths_key=None, block_table=None, layout_query='BSND', layout_key='BSND', sparse_count=2048, sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1) -> Tensor
24```25```
25 26 
26## 参数说明27## 参数说明
28+>
27> [!NOTE] 29> [!NOTE]
28>30>
29> - query、key、weights、query_dequant_scale、key_dequant_scale参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。31> - query、key、weights、query_dequant_scale、key_dequant_scale参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。
30> - 使用S1和S2分别表示query和key的输入样本序列长度,N1和N2分别表示query和key对应的多头数,k表示最后选取的索引个数。参数query中的D和参数key中的D值相等为128。T1和T2分别表示query和key的输入样本序列长度的累加和。32> - 使用S1和S2分别表示query和key的输入样本序列长度,N1和N2分别表示query和key对应的多头数,k表示最后选取的索引个数。参数query中的D和参数key中的D值相等为128。T1和T2分别表示query和key的输入样本序列长度的累加和。
31-- **query**(`Tensor`):必选参数,表示输入Index Query,对应公式中的$Q_{index}^{INT8}\in\R^{g\times d}$。不支持非连续,数据格式支持$ND$,数据类型支持`int8`。`layout_query`为BSND时shape为[B,S1,N1,D],当`layout_query`为TND时shape为[T1,N1,D],N1仅支持64。33+>
34+- **query**`Tensor`):必选参数,表示输入Index Query,对应公式中的$Q_{index}^{INT8}\in\R^{g\times d}$。不支持非连续,数据格式支持$ND$,数据类型支持`int8`。`layout_query`为BSND时shape为[B,S1,N1,D],当`layout_query`为TND时shape为[T1,N1,D],N1仅支持64。
32 35
33-- **key**(`Tensor`):必选参数,表示输入Index Key,对应公式中的$K_{index}^{INT8}\in\R^{S_{k}\times d}$。不支持非连续,数据格式支持$ND$,数据类型支持`int8`,layout\_key为PA_BSND时shape为[block\_count, block\_size, N2, D],其中block\_count为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的整数倍,最大支持到1024。`layout_kv`为BSND时shape为[B, S2, N2, D],`layout_kv`为TND时shape为[T2, N2, D],N2仅支持1。36+- **key**(`Tensor`):必选参数,表示输入Index Key,对应公式中的$K_{index}^{INT8}\in\R^{S_{k}\times d}$。不支持非连续,数据格式支持$ND$,数据类型支持`int8`,layout\_key为PA_BSND时shape为[block\_count, block\_size, N2, D],其中block\_count为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的整数倍,最大支持到1024。`layout_kv`为BSND时shape为[B, S2, N2, D],`layout_kv`为TND时shape为[T2, N2, D],N2仅支持1。
34 37
35-- **weights**(`Tensor`):必选参数,表示权重系数,对应公式中的$W$。不支持非连续,数据格式支持$ND$,数据类型支持`float16`,支持输入shape[B,S1,N1]、[T,N1]。38+- **weights**(`Tensor`):必选参数,表示权重系数,对应公式中的$W$。不支持非连续,数据格式支持$ND$,数据类型支持`float16`,支持输入shape[B,S1,N1]、[T,N1]。
36 39 
37-- **query_dequant_scale**(`Tensor`):必选参数,表示Index Query的反量化系数$Scale_Q$ 。不支持非连续,数据格式支持$ND$,数据类型支持`float16`,支持输入shape[B,S1,N1]、[T,N1]。40+- **query_dequant_scale**(`Tensor`):必选参数,表示Index Query的反量化系数$Scale_Q$ 。不支持非连续,数据格式支持$ND$,数据类型支持`float16`,支持输入shape[B,S1,N1]、[T,N1]。
38 41 
39-- **key_dequant_scale**(`Tensor`):必选参数,表示Index Key的反量化系数,对应公式中的$Scale_K^T$。不支持非连续,数据格式支持$ND$,数据类型支持`float16`,layout\_key为PA_BSND时shape为[block\_count, block\_size, N2],其中block\_count为PageAttention时block总数,block\_size为一个block的token数。42+- **key_dequant_scale**(`Tensor`):必选参数,表示Index Key的反量化系数,对应公式中的$Scale_K^T$。不支持非连续,数据格式支持$ND$,数据类型支持`float16`,layout\_key为PA_BSND时shape为[block\_count, block\_size, N2],其中block\_count为PageAttention时block总数,block\_size为一个block的token数。
40 43 
41-- **query\_quant\_mode**(`int`):可选参数,用于标识输入`query`的量化模式,当前仅支持Per-Token-Head量化模式,当前仅支持传入0。44+- **query\_quant\_mode**(`int`):可选参数,用于标识输入`query`的量化模式,当前仅支持Per-Token-Head量化模式,当前仅支持传入0。
42 45 
43-- **key\_quant\_mode**(`int`):可选参数,用于标识输入`key`的量化模式,当前仅支持Per-Token-Head量化模式,当前仅支持传入0。46+- **key\_quant\_mode**(`int`):可选参数,用于标识输入`key`的量化模式,当前仅支持Per-Token-Head量化模式,当前仅支持传入0。
44 47 
45- <strong>*</strong>:代表其之前的参数是位置相关的,必须按照顺序输入;之后的参数是可选参数,位置无关,不赋值会使用默认值。48- <strong>*</strong>:代表其之前的参数是位置相关的,必须按照顺序输入;之后的参数是可选参数,位置无关,不赋值会使用默认值。
46 49 
47-- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。该入参中每个Batch的有效token数不超过`query`中的维度S大小且不小于0。支持长度为B的一维tensor。当`layout_query`为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。不能出现负值。50+- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。该入参中每个Batch的有效token数不超过`query`中的维度S大小且不小于0。支持长度为B的一维tensor。当`layout_query`为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。不能出现负值。
48 51 
49-- **actual\_seq\_lengths\_key**(`Tensor`):可选参数,表示不同Batch中`key`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和key的shape的S长度相同。该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。当`layout_kv`为TND或PA_BSND时,该入参必须传入,`layout_kv`为TND,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。52+- **actual\_seq\_lengths\_key**(`Tensor`):可选参数,表示不同Batch中`key`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和key的shape的S长度相同。该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。当`layout_kv`为TND或PA_BSND时,该入参必须传入,`layout_kv`为TND,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。
50 53 
51-- **block\_table**(`Tensor`):可选参数,表示PageAttention中KV存储使用的block映射表,数据格式支持$ND$,数据类型支持`int32`。PageAttention场景下,block\_table必须为二维,第一维长度需要等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为每个batch中最大actual\_seq\_lengths\_key对应的block数量),支持block_size取值为16的整数倍,最大支持到1024。54+- **block\_table**(`Tensor`):可选参数,表示PageAttention中KV存储使用的block映射表,数据格式支持$ND$,数据类型支持`int32`。PageAttention场景下,block\_table必须为二维,第一维长度需要等于B,第二维长度不能小于maxBlockNumPerSeq(maxBlockNumPerSeq为每个batch中最大actual\_seq\_lengths\_key对应的block数量),支持block_size取值为16的整数倍,最大支持到1024。
52 55 
53-- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,当前支持BSND、TND,默认值"BSND"。56+- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,当前支持BSND、TND,默认值"BSND"。
54 57 
55-- **layout\_key**(`str`):可选参数,用于标识输入`key`的数据排布格式,当前支持PA_BSND、BSND、TND,默认值"BSND"。在非PageAttention场景下,layout\_key应与layout\_query保持一致。58+- **layout\_key**(`str`):可选参数,用于标识输入`key`的数据排布格式,当前支持PA_BSND、BSND、TND,默认值"BSND"。在非PageAttention场景下,layout\_key应与layout\_query保持一致。
56 59 
57-- **sparse\_count**(`int`):可选参数,代表topK阶段需要保留的block数量,支持[1, 2048],数据类型支持`int32`。60+- **sparse\_count**(`int`):可选参数,代表topK阶段需要保留的block数量,支持[1, 2048],数据类型支持`int32`。
58 61 
59-- **sparse\_mode**(`int`):可选参数,表示sparse的模式,支持0/3,数据类型支持`int32`。 sparse\_mode为0时,代表defaultMask模式。sparse\_mode为3时,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景。62+- **sparse\_mode**(`int`):可选参数,表示sparse的模式,支持0/3,数据类型支持`int32`。 sparse\_mode为0时,代表defaultMask模式。sparse\_mode为3时,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景。
60 63 
61-- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。64+- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。
62 65 
63-- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。66+- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。
64 67 
65## 返回值说明68## 返回值说明
69+ 
66`Tensor`70`Tensor`
67 71 
68代表公式中的输出Out。数据格式支持$ND$,数据类型支持`int32`,支持输出shape[B,S1,N2,k]或[T,N2,k]。72代表公式中的输出Out。数据格式支持$ND$,数据类型支持`int32`,支持输出shape[B,S1,N2,k]或[T,N2,k]。
69 73 
70## 约束说明74## 约束说明
71-- 该接口支持图模式。75+ 
72-- 该接口要求$W \odot Scale_Q$的结果在`float16`的表示范围内76+- 该接口支持图模式
73-- 该接口的TopK过程对NAN排序是未定义行为77+- 该接口要求$W \odot Scale_Q$结果在`float16`的表示范围内
78+- 该接口的TopK过程对NAN排序是未定义行为。
74 79 
75## 调用示例80## 调用示例
76 81 
77-- 单算子模式调用82+- 单算子模式调用
83+ 
78 ```python84 ```python
79 import torch85 import torch
80 import torch_npu86 import torch_npu
@@ -124,7 +130,8 @@ torch_npu.npu_quant_lightning_indexer(query, key, weights, query_dequant_scale,
124 layout_key=layout_key, sparse_count=sparse_count,130 layout_key=layout_key, sparse_count=sparse_count,
125 sparse_mode=sparse_mode)131 sparse_mode=sparse_mode)
126 ```132 ```
127-- 图模式调用133+ 
134+- 图模式调用
128 135 
129 ```python136 ```python
130 import torch137 import torch
@@ -195,4 +202,4 @@ torch_npu.npu_quant_lightning_indexer(query, key, weights, query_dequant_scale,
195 block_table=block_table, query_quant_mode=query_quant_mode,202 block_table=block_table, query_quant_mode=query_quant_mode,
196 key_quant_mode=key_quant_mode, layout_query=layout_query,203 key_quant_mode=key_quant_mode, layout_query=layout_query,
197 layout_key=layout_key, sparse_count=sparse_count, sparse_mode=sparse_mode)204 layout_key=layout_key, sparse_count=sparse_count, sparse_mode=sparse_mode)
198- ```205+ ```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_quant_matmul.md+4-3
@@ -28,7 +28,7 @@
28 28 
29## 函数原型29## 函数原型
30 30 
31-```31+```python
32torch_npu.npu_quant_matmul(x1, x2, scale, *, offset=None, pertoken_scale=None, bias=None, output_dtype=None, group_sizes=None) -> Tensor32torch_npu.npu_quant_matmul(x1, x2, scale, *, offset=None, pertoken_scale=None, bias=None, output_dtype=None, group_sizes=None) -> Tensor
33```33```
34 34 
@@ -69,6 +69,7 @@ torch_npu.npu_quant_matmul(x1, x2, scale, *, offset=None, pertoken_scale=None, b
69`Tensor`69`Tensor`
70 70 
71代表量化matmul的计算结果。71代表量化matmul的计算结果。
72+ 
72- 如果`output_dtype``float16`,输出的数据类型为`float16`73- 如果`output_dtype``float16`,输出的数据类型为`float16`
73- 如果`output_dtype``int8`或者`None`,输出的数据类型为`int8`74- 如果`output_dtype``int8`或者`None`,输出的数据类型为`int8`
74- 如果`output_dtype``bfloat16`,输出的数据类型为`bfloat16`75- 如果`output_dtype``bfloat16`,输出的数据类型为`bfloat16`
@@ -98,6 +99,7 @@ torch_npu.npu_quant_matmul(x1, x2, scale, *, offset=None, pertoken_scale=None, b
98- 输入参数间支持的数据类型组合情况如下:99- 输入参数间支持的数据类型组合情况如下:
99 100 
100 **表 1** <term>Atlas 推理系列加速卡产品</term>101 **表 1** <term>Atlas 推理系列加速卡产品</term>
102+ 
101 |x1|x2|scale|offset|bias|pertoken_scale|output_dtype|103 |x1|x2|scale|offset|bias|pertoken_scale|output_dtype|
102 |---------|--------|--------|--------|--------|--------|--------|104 |---------|--------|--------|--------|--------|--------|--------|
103 |int8|int8|int64/float32|None|int32/None|None|float16|105 |int8|int8|int64/float32|None|int32/None|None|float16|
@@ -115,7 +117,6 @@ torch_npu.npu_quant_matmul(x1, x2, scale, *, offset=None, pertoken_scale=None, b
115 |int32|int32|float32|float16|None|float32|bfloat16/float16|117 |int32|int32|float32|float16|None|float32|bfloat16/float16|
116 |int8|int8|float32/bfloat16|None|int32/None|None|int32|118 |int8|int8|float32/bfloat16|None|int32/None|None|int32|
117 119 
118- 
119## 调用示例120## 调用示例
120 121 
121- 单算子调用122- 单算子调用
@@ -450,4 +451,4 @@ torch_npu.npu_quant_matmul(x1, x2, scale, *, offset=None, pertoken_scale=None, b
450 0.0000e+00, -1.0000e+00],451 0.0000e+00, -1.0000e+00],
451 [ 0.0000e+00, -1.0000e+00, -1.0000e+00, ..., -1.0000e+00,452 [ 0.0000e+00, -1.0000e+00, -1.0000e+00, ..., -1.0000e+00,
452 0.0000e+00, -1.0000e+00]], device='npu:0', dtype=torch.bfloat16)453 0.0000e+00, -1.0000e+00]], device='npu:0', dtype=torch.bfloat16)
453- ```454+ ```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_quant_matmul_gelu.md+2-2
@@ -26,7 +26,6 @@
26 qbmmout = x1@x2 * x2Scale * x1Scale26 qbmmout = x1@x2 * x2Scale * x1Scale
27 $$27 $$
28 28 
29-
30 - GELU激活函数,GELU类型由`approximate`输入指定:29 - GELU激活函数,GELU类型由`approximate`输入指定:
31 -`approximate``gelu_tanh`时:30 -`approximate``gelu_tanh`时:
32 $$31 $$
@@ -39,7 +38,7 @@
39 38 
40## 函数原型39## 函数原型
41 40 
42-```41+```python
43torch_npu.npu_quant_matmul_gelu(x1, x2, x1_scale, x2_scale, *, bias=None, approximate="gelu_erf") -> Tensor42torch_npu.npu_quant_matmul_gelu(x1, x2, x1_scale, x2_scale, *, bias=None, approximate="gelu_erf") -> Tensor
44```43```
45 44 
@@ -68,6 +67,7 @@ A8W8量化场景下,支持昇腾亲和的$NZ$数据排布格式,可通过`to
68`Tensor`67`Tensor`
69 68 
70代表量化矩阵乘融合GELU激活的计算结果。69代表量化矩阵乘融合GELU激活的计算结果。
70+ 
71- 输出数据类型的确定规则:71- 输出数据类型的确定规则:
72 - 如果`x2_scale`的数据类型为`float32`,输出的数据类型为`float16`72 - 如果`x2_scale`的数据类型为`float32`,输出的数据类型为`float16`
73 - 如果`x2_scale`的数据类型为`bfloat16`,输出的数据类型为`bfloat16`73 - 如果`x2_scale`的数据类型为`bfloat16`,输出的数据类型为`bfloat16`
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_quant_matmul_reduce_sum.md+5-5
@@ -19,11 +19,10 @@ $$
19 19 
20## 函数原型20## 函数原型
21 21 
22-```22+```python
23torch_npu.npu_quant_matmul_reduce_sum(x1, x2, *, x1_scale=None, x2_scale=None) -> Tensor23torch_npu.npu_quant_matmul_reduce_sum(x1, x2, *, x1_scale=None, x2_scale=None) -> Tensor
24```24```
25 25 
26- 
27## 参数说明26## 参数说明
28 27 
29- **x1** (`Tensor`):必选参数,数据类型支持`int8`,数据格式支持$ND$,shape支持3维,形状为(batch, m, k)。28- **x1** (`Tensor`):必选参数,数据类型支持`int8`,数据格式支持$ND$,shape支持3维,形状为(batch, m, k)。
@@ -37,7 +36,6 @@ torch_npu.npu_quant_matmul_reduce_sum(x1, x2, *, x1_scale=None, x2_scale=None) -
37- **x2_scale** (`Tensor`):必选关键字参数,对应公式中的$x2Scale$。数据类型支持`bfloat16`,数据格式支持$ND$,shape支持1维,形状为(n,)。36- **x2_scale** (`Tensor`):必选关键字参数,对应公式中的$x2Scale$。数据类型支持`bfloat16`,数据格式支持$ND$,shape支持1维,形状为(n,)。
38 - 在实际计算时,`x2_scale`会被广播到(batch,m,n)。37 - 在实际计算时,`x2_scale`会被广播到(batch,m,n)。
39 38 
40- 
41## 返回值说明39## 返回值说明
42 40 
43`Tensor`41`Tensor`
@@ -50,15 +48,16 @@ torch_npu.npu_quant_matmul_reduce_sum(x1, x2, *, x1_scale=None, x2_scale=None) -
50- 该接口支持静态图模式。48- 该接口支持静态图模式。
51- 传入的`x1``x2``x1_scale``x2_scale`不能是空。49- 传入的`x1``x2``x1_scale``x2_scale`不能是空。
52- 输入和输出支持以下数据类型组合:50- 输入和输出支持以下数据类型组合:
51+ 
53 | x1 | x2 | x1_scale | x2_scale | out |52 | x1 | x2 | x1_scale | x2_scale | out |
54 |------|------|---------|----------|----------|53 |------|------|---------|----------|----------|
55 | int8 | int8 | float32 | bfloat16 | bfloat16 |54 | int8 | int8 | float32 | bfloat16 | bfloat16 |
56 55 
57- 
58## 调用示例56## 调用示例
59 57 
60- 单算子调用58- 单算子调用
61- ```59+ 
60+ ```python
62 import torch61 import torch
63 import torch_npu62 import torch_npu
64 63 
@@ -72,6 +71,7 @@ torch_npu.npu_quant_matmul_reduce_sum(x1, x2, *, x1_scale=None, x2_scale=None) -
72 ```71 ```
73 72 
74- 图模式调用73- 图模式调用
74+ 
75 ```python75 ```python
76 import torch76 import torch
77 import torch_npu77 import torch_npu
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_quant_scatter.md+2-2
@@ -13,7 +13,7 @@
13 13 
14## 函数原型14## 函数原型
15 15 
16-```16+```python
17torch_npu.npu_quant_scatter(input, indices, updates, quant_scales, quant_zero_points=None, axis=0, quant_axis=1, reduce='update', int? dst_type=None, str? round_mode='rint') -> Tensor17torch_npu.npu_quant_scatter(input, indices, updates, quant_scales, quant_zero_points=None, axis=0, quant_axis=1, reduce='update', int? dst_type=None, str? round_mode='rint') -> Tensor
18```18```
19 19 
@@ -39,6 +39,7 @@ torch_npu.npu_quant_scatter(input, indices, updates, quant_scales, quant_zero_po
39- **reduce** (`str`):可选参数,表示数据操作方式;当前只支持`'update'`,即更新操作。39- **reduce** (`str`):可选参数,表示数据操作方式;当前只支持`'update'`,即更新操作。
40 40 
41## 返回值说明41## 返回值说明
42+ 
42`Tensor`43`Tensor`
43 44 
44代表`input`被更新后的结果。45代表`input`被更新后的结果。
@@ -258,4 +259,3 @@ torch_npu.npu_quant_scatter(input, indices, updates, quant_scales, quant_zero_po
258 1, 0, 0, 2, 1, 0, 1, 1, 0, 1, 0, 1, 0, 1, 0]]],259 1, 0, 0, 2, 1, 0, 1, 1, 0, 1, 0, 1, 0, 1, 0]]],
259 device='npu:0', dtype=torch.int8) torch.Size([11, 1, 32])260 device='npu:0', dtype=torch.int8) torch.Size([11, 1, 32])
260 ```261 ```
261- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_quant_scatter_.md+2-2
@@ -13,7 +13,7 @@
13 13 
14## 函数原型14## 函数原型
15 15 
16-```16+```python
17torch_npu.npu_quant_scatter_(input, indices, updates, quant_scales, quant_zero_points=None, axis=0, quant_axis=1, reduce='update', int? dst_type=None, str? round_mode='rint') -> Tensor17torch_npu.npu_quant_scatter_(input, indices, updates, quant_scales, quant_zero_points=None, axis=0, quant_axis=1, reduce='update', int? dst_type=None, str? round_mode='rint') -> Tensor
18```18```
19 19 
@@ -39,6 +39,7 @@ torch_npu.npu_quant_scatter_(input, indices, updates, quant_scales, quant_zero_p
39- **reduce** (`str`):可选参数,表示数据操作方式;当前只支持`update`,即更新操作。39- **reduce** (`str`):可选参数,表示数据操作方式;当前只支持`update`,即更新操作。
40 40 
41## 返回值说明41## 返回值说明
42+ 
42`Tensor`43`Tensor`
43 44 
44代表`input`被更新后的结果。45代表`input`被更新后的结果。
@@ -259,4 +260,3 @@ torch_npu.npu_quant_scatter_(input, indices, updates, quant_scales, quant_zero_p
259 1, 1, 0, 1, 1, 0, 0, 0, 0, 0, 0, 1, 0, 1, 1]]],260 1, 1, 0, 1, 1, 0, 0, 0, 0, 0, 0, 1, 0, 1, 1]]],
260 device='npu:0', dtype=torch.int8) torch.Size([11, 1, 32])261 device='npu:0', dtype=torch.int8) torch.Size([11, 1, 32])
261 ```262 ```
262- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_quantize.md+2-2
@@ -25,7 +25,7 @@
25 25 
26## 函数原型26## 函数原型
27 27 
28-```28+```python
29torch_npu.npu_quantize(input, scales, zero_points, dtype, axis=1, div_mode=True) -> Tensor29torch_npu.npu_quantize(input, scales, zero_points, dtype, axis=1, div_mode=True) -> Tensor
30```30```
31 31 
@@ -67,6 +67,7 @@ torch_npu.npu_quantize(input, scales, zero_points, dtype, axis=1, div_mode=True)
67- **div_mode** (`bool`):可选参数,表示计算`scales`模式,对应公式中的`div_mode`。当`div_mode`为`True`时,表示用除法计算`scales`;`div_mode`为`False`时,表示用乘法计算`scales`,默认值为`True`。67- **div_mode** (`bool`):可选参数,表示计算`scales`模式,对应公式中的`div_mode`。当`div_mode`为`True`时,表示用除法计算`scales`;`div_mode`为`False`时,表示用乘法计算`scales`,默认值为`True`。
68 68 
69## 返回值说明69## 返回值说明
70+ 
70`Tensor`71`Tensor`
71 72 
72对应公式中的`result`。数据类型由参数`dtype`指定,如果参数`dtype``quint4x2`,输出的`dtype``int32`,shape的最后一维是`input`的shape最后一维的1/8,shape其他维度和`input`的shape其他维度保持一致;如果参数`dtype`不为`quint4x2`时,shape与输入`input`的shape保持一致。输出的数据格式与输入`input`的数据格式保持一致,且当数据格式为$NZ$时,数据类型仅支持INT32。支持空Tensor,支持非连续的Tensor。73对应公式中的`result`。数据类型由参数`dtype`指定,如果参数`dtype``quint4x2`,输出的`dtype``int32`,shape的最后一维是`input`的shape最后一维的1/8,shape其他维度和`input`的shape其他维度保持一致;如果参数`dtype`不为`quint4x2`时,shape与输入`input`的shape保持一致。输出的数据格式与输入`input`的数据格式保持一致,且当数据格式为$NZ$时,数据类型仅支持INT32。支持空Tensor,支持非连续的Tensor。
@@ -205,4 +206,3 @@ torch_npu.npu_quantize(input, scales, zero_points, dtype, axis=1, div_mode=True)
205 dtype=torch.int8)206 dtype=torch.int8)
206 207 
207 ```208 ```
208- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_recurrent_gated_delta_rule.md+4-1
@@ -26,7 +26,7 @@
26 26 
27## 函数原型27## 函数原型
28 28 
29-```29+```python
30torch_npu.npu_recurrent_gated_delta_rule(query, key, value, state, *, beta=None, scale=None, actual_seq_lengths=None, ssm_state_indices=None, num_accepted_tokens=None, g=None, gk=None) -> Tensor30torch_npu.npu_recurrent_gated_delta_rule(query, key, value, state, *, beta=None, scale=None, actual_seq_lengths=None, ssm_state_indices=None, num_accepted_tokens=None, g=None, gk=None) -> Tensor
31```31```
32 32 
@@ -61,6 +61,7 @@ torch_npu.npu_recurrent_gated_delta_rule(query, key, value, state, *, beta=None,
61公式中的$o$,注意力计算结果。输出的数据类型为`bfloat16`,数据格式为ND,shape为($T$, $N_v$, $D_v$)。61公式中的$o$,注意力计算结果。输出的数据类型为`bfloat16`,数据格式为ND,shape为($T$, $N_v$, $D_v$)。
62 62 
63## 约束说明63## 约束说明
64+ 
64- 参数里Shape使用的变量如下:65- 参数里Shape使用的变量如下:
65 - $T=\sum_i^B L_i$ 表示累积序列长度。66 - $T=\sum_i^B L_i$ 表示累积序列长度。
66 - $B$ 表示batch size。67 - $B$ 表示batch size。
@@ -76,6 +77,7 @@ torch_npu.npu_recurrent_gated_delta_rule(query, key, value, state, *, beta=None,
76## 调用示例77## 调用示例
77 78 
78- 单算子调用79- 单算子调用
80+ 
79 ```python81 ```python
80 import torch82 import torch
81 import torch_npu83 import torch_npu
@@ -107,6 +109,7 @@ torch_npu.npu_recurrent_gated_delta_rule(query, key, value, state, *, beta=None,
107 ```109 ```
108 110 
109- 静态图模式调用111- 静态图模式调用
112+ 
110 ```python113 ```python
111 import torch114 import torch
112 import torch_npu115 import torch_npu
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_rms_norm_quant.md+4-2
@@ -25,7 +25,7 @@
25 25 
26## 函数原型26## 函数原型
27 27 
28-```28+```python
29torch_npu.npu_rms_norm_quant(x, gamma, beta, scale, offset, epsilon=1e-06) -> Tensor29torch_npu.npu_rms_norm_quant(x, gamma, beta, scale, offset, epsilon=1e-06) -> Tensor
30```30```
31 31 
@@ -52,6 +52,7 @@ torch_npu.npu_rms_norm_quant(x, gamma, beta, scale, offset, epsilon=1e-06) -> Te
52- **epsilon** (`float`):可选参数,对应公式中的$eps$,用于防止除零错误,默认值为 `1e-6`。建议传入较小的正数。52- **epsilon** (`float`):可选参数,对应公式中的$eps$,用于防止除零错误,默认值为 `1e-6`。建议传入较小的正数。
53 53 
54- **dst_dtype** (`int`): 可选参数,指定量化输出的类型,传`None`时当做int8处理,支持取值`int8`、`quint4x2`。54- **dst_dtype** (`int`): 可选参数,指定量化输出的类型,传`None`时当做int8处理,支持取值`int8`、`quint4x2`。
55+ 
55## 返回值说明56## 返回值说明
56 57
57 `Tensor`58 `Tensor`
@@ -81,6 +82,7 @@ torch_npu.npu_rms_norm_quant(x, gamma, beta, scale, offset, epsilon=1e-06) -> Te
81 | float16 | float16 | float16 | float16 | int8 | double |int32 |82 | float16 | float16 | float16 | float16 | int8 | double |int32 |
82 83 
83## 调用示例84## 调用示例
85+ 
84```python86```python
85>>> import torch87>>> import torch
86>>> import torch_npu88>>> import torch_npu
@@ -94,4 +96,4 @@ torch_npu.npu_rms_norm_quant(x, gamma, beta, scale, offset, epsilon=1e-06) -> Te
94>>> y.cpu().numpy()96>>> y.cpu().numpy()
95 tensor([ 1, -1, 2, 0, -2, 1, 0, 1, 2, 0, 2, 0, 0, 0, 0, 0],97 tensor([ 1, -1, 2, 0, -2, 1, 0, 1, 2, 0, 2, 0, 0, 0, 0, 0],
96 device='npu:0', dtype=torch.int8)98 device='npu:0', dtype=torch.int8)
97-```99+```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_rotary_mul.md+11-5
@@ -11,8 +11,8 @@
11 11 
12## 功能说明12## 功能说明
13 13 
14-- API功能:实现Rotary Position Embedding (RoPE) 旋转位置编码,通过对输入特征进行二维平面旋转注入位置信息。14+- API功能:实现Rotary Position Embedding (RoPE) 旋转位置编码,通过对输入特征进行二维平面旋转注入位置信息。
15-- 计算公式:15+- 计算公式:
16 $$16 $$
17 output = x * cos + rotate(x) * sin17 output = x * cos + rotate(x) * sin
18 $$18 $$
@@ -34,7 +34,7 @@
34 rotate(x) = x \cdot rotate\\34 rotate(x) = x \cdot rotate\\
35 $$35 $$
36 36 
37-- 等价计算逻辑:37+- 等价计算逻辑:
38 38
39 可使用`fused_rotary_position_embedding`等价替换`torch_npu.npu_rotary_mul`,两者计算逻辑一致。39 可使用`fused_rotary_position_embedding`等价替换`torch_npu.npu_rotary_mul`,两者计算逻辑一致。
40 40
@@ -62,7 +62,7 @@
62 62 
63## 函数原型63## 函数原型
64 64 
65-```65+```python
66torch_npu.npu_rotary_mul(input, r1, r2, rotary_mode='half', rotate=None) -> Tensor66torch_npu.npu_rotary_mul(input, r1, r2, rotary_mode='half', rotate=None) -> Tensor
67```67```
68 68 
@@ -78,6 +78,7 @@ torch_npu.npu_rotary_mul(input, r1, r2, rotary_mode='half', rotate=None) -> Tens
78- **rotate** (`Tensor`):可选参数,表示实现`input`位置变换的等价变化矩阵,输入维度支持2维,数据类型支持`float16``bfloat16``float32`,构造方式参考调用示例,默认值为None。78- **rotate** (`Tensor`):可选参数,表示实现`input`位置变换的等价变化矩阵,输入维度支持2维,数据类型支持`float16``bfloat16``float32`,构造方式参考调用示例,默认值为None。
79 79 
80## 返回值说明80## 返回值说明
81+ 
81`Tensor`82`Tensor`
82 83 
83输出计算结果,shape和dtype需与`input`一致。84输出计算结果,shape和dtype需与`input`一致。
@@ -125,9 +126,11 @@ torch_npu.npu_rotary_mul(input, r1, r2, rotary_mode='half', rotate=None) -> Tens
125 那么可以构造一个rotate矩阵,实现调用一次完成x的旋转位置编码计算功能,rotate矩阵构造如下:126 那么可以构造一个rotate矩阵,实现调用一次完成x的旋转位置编码计算功能,rotate矩阵构造如下:
126 $$rotate = diag(rotate1, rotate2, rotate3) = \begin{pmatrix}rotate1&0&0\\0&rotate2&0\\0&0&rotate3\\\end{pmatrix}$$127 $$rotate = diag(rotate1, rotate2, rotate3) = \begin{pmatrix}rotate1&0&0\\0&rotate2&0\\0&0&rotate3\\\end{pmatrix}$$
127 其中rotate1、rotate2、rotate3分别为x1、x2、x3的旋转编码矩阵,单个旋转矩阵构建参考调用示例。128 其中rotate1、rotate2、rotate3分别为x1、x2、x3的旋转编码矩阵,单个旋转矩阵构建参考调用示例。
129+ 
128## 调用示例130## 调用示例
129 131 
130- 四维输入示例:132- 四维输入示例:
133+ 
131 ```python134 ```python
132 >>> import torch135 >>> import torch
133 >>> import torch_npu136 >>> import torch_npu
@@ -165,7 +168,9 @@ torch_npu.npu_rotary_mul(input, r1, r2, rotary_mode='half', rotate=None) -> Tens
165 [-0.2269, -0.1447, -0.0395, ..., 0.1374, 0.2142, 0.3628]]]],168 [-0.2269, -0.1447, -0.0395, ..., 0.1374, 0.2142, 0.3628]]]],
166 device='npu:0')169 device='npu:0')
167 ```170 ```
171+ 
168- 三维输入示例:172- 三维输入示例:
173+ 
169 ```python174 ```python
170 >>> import torch175 >>> import torch
171 >>> import torch_npu176 >>> import torch_npu
@@ -190,6 +195,7 @@ torch_npu.npu_rotary_mul(input, r1, r2, rotary_mode='half', rotate=None) -> Tens
190 ```195 ```
191 196 
192- rotate生成示例:197- rotate生成示例:
198+ 
193 ```python199 ```python
194 import torch200 import torch
195 import torch_npu201 import torch_npu
@@ -230,4 +236,4 @@ torch_npu.npu_rotary_mul(input, r1, r2, rotary_mode='half', rotate=None) -> Tens
230 r1 = torch.rand(1, 2, 1, 128).npu()236 r1 = torch.rand(1, 2, 1, 128).npu()
231 r2 = torch.rand(1, 2, 1, 128).npu()237 r2 = torch.rand(1, 2, 1, 128).npu()
232 out = torch_npu.npu_rotary_mul(x, r1, r2,"interleave", inter_mat_128.npu())238 out = torch_npu.npu_rotary_mul(x, r1, r2,"interleave", inter_mat_128.npu())
233- ```239+ ```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_scaled_masked_softmax.md+2-2
@@ -15,7 +15,7 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.npu_scaled_masked_softmax(x, mask, scale=1.0, fixed_triu_mask=False) -> Tensor19torch_npu.npu_scaled_masked_softmax(x, mask, scale=1.0, fixed_triu_mask=False) -> Tensor
20```20```
21 21 
@@ -27,6 +27,7 @@ torch_npu.npu_scaled_masked_softmax(x, mask, scale=1.0, fixed_triu_mask=False) -
27- **fixed_triu_mask**`bool`):预留参数,功能未完成,默认值为`False`,当前只支持`False`。该功能完成后可支持自动生成上三角`bool`掩码。27- **fixed_triu_mask**`bool`):预留参数,功能未完成,默认值为`False`,当前只支持`False`。该功能完成后可支持自动生成上三角`bool`掩码。
28 28 
29## 返回值说明29## 返回值说明
30+ 
30`Tensor`31`Tensor`
31 32 
32一个`Tensor`类型的输出,输入`x`经过`mask`后在最后一维的`Softmax`结果,输出shape与`x`一致。支持数据类型:`float16``float32``bfloat16`。支持格式:$[ND,FRACTAL\_NZ]$。33一个`Tensor`类型的输出,输入`x`经过`mask`后在最后一维的`Softmax`结果,输出shape与`x`一致。支持数据类型:`float16``float32``bfloat16`。支持格式:$[ND,FRACTAL\_NZ]$。
@@ -50,4 +51,3 @@ torch_npu.npu_scaled_masked_softmax(x, mask, scale=1.0, fixed_triu_mask=False) -
50>>> output.shape51>>> output.shape
51torch.size([4, 4, 2048, 2048])52torch.size([4, 4, 2048, 2048])
52```53```
53- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_scatter_nd_update.md+3-5
@@ -15,7 +15,7 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.npu_scatter_nd_update(input, indices, updates) -> Tensor19torch_npu.npu_scatter_nd_update(input, indices, updates) -> Tensor
20```20```
21 21 
@@ -26,9 +26,7 @@ torch_npu.npu_scatter_nd_update(input, indices, updates) -> Tensor
26 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float32``float16``bool``bfloat16``int64``int8`26 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float32``float16``bool``bfloat16``int64``int8`
27 - <term>Atlas 推理系列加速卡产品</term>:数据类型支持`float32``float16``bool`27 - <term>Atlas 推理系列加速卡产品</term>:数据类型支持`float32``float16``bool`
28 - <term>Atlas 训练系列产品</term>:数据类型支持`float32``float16``bool`28 - <term>Atlas 训练系列产品</term>:数据类型支持`float32``float16``bool`
29- 29+
30-
31- 
32- **indices** (`Tensor`):必选输入,索引张量,数据类型支持`int32``int64`,数据格式支持$ND$,支持非连续的Tensor,`indices`中的索引数据不支持越界。30- **indices** (`Tensor`):必选输入,索引张量,数据类型支持`int32``int64`,数据格式支持$ND$,支持非连续的Tensor,`indices`中的索引数据不支持越界。
33- **updates** (`Tensor`):必选输入,更新数据张量,数据格式支持$ND$,支持非连续的Tensor,数据类型需要与`input`一致。31- **updates** (`Tensor`):必选输入,更新数据张量,数据格式支持$ND$,支持非连续的Tensor,数据类型需要与`input`一致。
34 - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`float32``float16``bool``bfloat16``int64``int8`32 - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`float32``float16``bool``bfloat16``int64``int8`
@@ -37,6 +35,7 @@ torch_npu.npu_scatter_nd_update(input, indices, updates) -> Tensor
37 - <term>Atlas 训练系列产品</term>:数据类型支持`float32``float16``bool`35 - <term>Atlas 训练系列产品</term>:数据类型支持`float32``float16``bool`
38 36 
39## 返回值37## 返回值
38+ 
40`Tensor`39`Tensor`
41 40 
42代表`input`被更新后的结果。41代表`input`被更新后的结果。
@@ -137,4 +136,3 @@ torch_npu.npu_scatter_nd_update(input, indices, updates) -> Tensor
137 [0.2910, 0.2468, 0.5488, 0.9761, 0.9785]], device='npu:0',136 [0.2910, 0.2468, 0.5488, 0.9761, 0.9785]], device='npu:0',
138 dtype=torch.float16)137 dtype=torch.float16)
139 ```138 ```
140- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_scatter_nd_update_.md+3-5
@@ -15,7 +15,7 @@
15 15 
16## 函数原型16## 函数原型
17 17 
18-```18+```python
19torch_npu.npu_scatter_nd_update_(input, indices, updates) -> Tensor19torch_npu.npu_scatter_nd_update_(input, indices, updates) -> Tensor
20```20```
21 21 
@@ -26,9 +26,7 @@ torch_npu.npu_scatter_nd_update_(input, indices, updates) -> Tensor
26 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float32``float16``bool``bfloat16``int64``int8`26 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:数据类型支持`float32``float16``bool``bfloat16``int64``int8`
27 - <term>Atlas 推理系列加速卡产品</term>:数据类型支持`float32``float16``bool`27 - <term>Atlas 推理系列加速卡产品</term>:数据类型支持`float32``float16``bool`
28 - <term>Atlas 训练系列产品</term>:数据类型支持`float32``float16``bool`28 - <term>Atlas 训练系列产品</term>:数据类型支持`float32``float16``bool`
29- 29+
30-
31- 
32- **indices** (`Tensor`):必选输入,索引张量,数据类型支持`int32``int64`,数据格式支持$ND$,支持非连续的Tensor,`indices`中的索引数据不支持越界。30- **indices** (`Tensor`):必选输入,索引张量,数据类型支持`int32``int64`,数据格式支持$ND$,支持非连续的Tensor,`indices`中的索引数据不支持越界。
33- **updates** (`Tensor`):必选输入,更新数据张量,数据格式支持$ND$,支持非连续的Tensor,数据类型需要与`input`一致。31- **updates** (`Tensor`):必选输入,更新数据张量,数据格式支持$ND$,支持非连续的Tensor,数据类型需要与`input`一致。
34 - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`float32``float16``bool``bfloat16``int64``int8`32 - <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:数据类型支持`float32``float16``bool``bfloat16``int64``int8`
@@ -37,6 +35,7 @@ torch_npu.npu_scatter_nd_update_(input, indices, updates) -> Tensor
37 - <term>Atlas 训练系列产品</term>:数据类型支持`float32``float16``bool`35 - <term>Atlas 训练系列产品</term>:数据类型支持`float32``float16``bool`
38 36
39## 返回值37## 返回值
38+ 
40`Tensor`39`Tensor`
41 40 
42代表`input`被更新后的结果。41代表`input`被更新后的结果。
@@ -127,4 +126,3 @@ torch_npu.npu_scatter_nd_update_(input, indices, updates) -> Tensor
127 # 执行上述代码的输出类似如下126 # 执行上述代码的输出类似如下
128 torch.Size([33, 5]) torch.float16127 torch.Size([33, 5]) torch.float16
129 ```128 ```
130- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_scatter_pa_kv_cache.md+3-1
@@ -14,6 +14,7 @@
14输入输出支持以下场景:14输入输出支持以下场景:
15 15 
16- 场景一:16- 场景一:
17+ 
17 ```python18 ```python
18 key:[batch, num_head, k_head_size]19 key:[batch, num_head, k_head_size]
19 value:[batch, num_head, v_head_size]20 value:[batch, num_head, v_head_size]
@@ -23,6 +24,7 @@
23 ``` 24 ```
24 25 
25- 场景二:26- 场景二:
27+ 
26 ```python 28 ```python
27 key:[batch, seq_len, num_head, k_head_size]29 key:[batch, seq_len, num_head, k_head_size]
28 value:[batch, seq_len, num_head, v_head_size]30 value:[batch, seq_len, num_head, v_head_size]
@@ -38,7 +40,7 @@
38 40 
39## 函数原型41## 函数原型
40 42 
41-```43+```python
42torch_npu.npu_scatter_pa_kv_cache(key, value, key_cache, value_cache, slot_mapping, *, compress_lens=None, compress_seq_offsets=None, seq_lens=None) -> ()44torch_npu.npu_scatter_pa_kv_cache(key, value, key_cache, value_cache, slot_mapping, *, compress_lens=None, compress_seq_offsets=None, seq_lens=None) -> ()
43```45```
44 46 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_sim_exponential_.md+2-8
@@ -1,6 +1,5 @@
1# torch_npu.npu_sim_exponential_1# torch_npu.npu_sim_exponential_
2 2 
3- 
4## 产品支持情况3## 产品支持情况
5 4 
6| 产品 | 是否支持 |5| 产品 | 是否支持 |
@@ -8,21 +7,18 @@
8|<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |7|<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | √ |
9|<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |8|<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | √ |
10 9 
11- 
12## 功能说明10## 功能说明
13 11 
14- API功能:根据参数`lambd`生成指数分布随机数,并原地填充至输入张量`input`12- API功能:根据参数`lambd`生成指数分布随机数,并原地填充至输入张量`input`
15- 计算公式:13- 计算公式:
16 $$f(x) = -1/λ * ln(1-u), u ~ Uniform(0, 1]$$14 $$f(x) = -1/λ * ln(1-u), u ~ Uniform(0, 1]$$
17 15 
18- 
19## 函数原型16## 函数原型
20 17 
21-```18+```python
22torch_npu.npu_sim_exponential_(input, lambd=1, *, generator=None) -> Tensor19torch_npu.npu_sim_exponential_(input, lambd=1, *, generator=None) -> Tensor
23```20```
24 21 
25- 
26## 参数说明22## 参数说明
27 23 
28**input**(`Tensor`):必选参数,源数据张量,公式中的$f(x)$。要求为连续的Tensor,数据类型支持`bfloat16``float16``float32`,数据格式支持$ND$,shape支持0~8维。24**input**(`Tensor`):必选参数,源数据张量,公式中的$f(x)$。要求为连续的Tensor,数据类型支持`bfloat16``float16``float32`,数据格式支持$ND$,shape支持0~8维。
@@ -31,14 +27,12 @@ torch_npu.npu_sim_exponential_(input, lambd=1, *, generator=None) -> Tensor
31 27 
32**generator**(`Generator`):可选参数,用于生成seed和offset,供aclnnSimThreadExponential算子使用,默认为None。28**generator**(`Generator`):可选参数,用于生成seed和offset,供aclnnSimThreadExponential算子使用,默认为None。
33 29 
34- 
35## 返回值说明30## 返回值说明
36 31 
37`Tensor`32`Tensor`
38 33 
39表示公式中的$f(x)$,即原地更新后的`input`张量。34表示公式中的$f(x)$,即原地更新后的`input`张量。
40 35 
41- 
42## 调用示例36## 调用示例
43 37 
44```python38```python
@@ -51,4 +45,4 @@ torch_npu.npu_sim_exponential_(input, lambd=1, *, generator=None) -> Tensor
51>>> input = torch.zeros(shape, dtype=torch.float32).npu()45>>> input = torch.zeros(shape, dtype=torch.float32).npu()
52>>> torch_npu.npu_sim_exponential_(input, lambd=1, generator=gen)46>>> torch_npu.npu_sim_exponential_(input, lambd=1, generator=gen)
53 47 
54-```48+```
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_sparse_flash_attention.md+38-33
@@ -1,12 +1,14 @@
1# torch_npu.npu_sparse_flash_attention1# torch_npu.npu_sparse_flash_attention
2 2 
3## 产品支持情况3## 产品支持情况
4+ 
4| 产品 | 是否支持 |5| 产品 | 是否支持 |
5| ------------------------------------------------------------ | :------: |6| ------------------------------------------------------------ | :------: |
6|<term>Atlas A2 推理系列产品</term> | √ |7|<term>Atlas A2 推理系列产品</term> | √ |
7|<term>Atlas A3 推理系列产品</term> | √ |8|<term>Atlas A3 推理系列产品</term> | √ |
8 9 
9## 功能说明10## 功能说明
11+ 
10- API功能:sparse_flash_attention(SFA)是针对大序列长度推理场景的高效注意力计算模块,该模块通过“只计算关键部分”大幅减少计算量,然而会引入大量的离散访存,造成数据搬运时间增加,进而影响整体性能。12- API功能:sparse_flash_attention(SFA)是针对大序列长度推理场景的高效注意力计算模块,该模块通过“只计算关键部分”大幅减少计算量,然而会引入大量的离散访存,造成数据搬运时间增加,进而影响整体性能。
11 13 
12- 计算公式:14- 计算公式:
@@ -20,72 +22,75 @@
20 22 
21## 函数原型23## 函数原型
22 24 
23-```25+```python
24torch_npu.npu_sparse_flash_attention(query, key, value, sparse_indices, scale_value, *, block_table=None, actual_seq_lengths_query=None, actual_seq_lengths_kv=None, query_rope=None, key_rope=None, sparse_block_size=1, layout_query='BSND', layout_kv='BSND', sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1, attention_mode=0, return_softmax_lse=False) -> (Tensor, Tensor, Tensor)26torch_npu.npu_sparse_flash_attention(query, key, value, sparse_indices, scale_value, *, block_table=None, actual_seq_lengths_query=None, actual_seq_lengths_kv=None, query_rope=None, key_rope=None, sparse_block_size=1, layout_query='BSND', layout_kv='BSND', sparse_mode=3, pre_tokens=2^63-1, next_tokens=2^63-1, attention_mode=0, return_softmax_lse=False) -> (Tensor, Tensor, Tensor)
25```27```
26 28 
27## 参数说明29## 参数说明
28 30 
29> [!NOTE] 31> [!NOTE]
32+>
30>- query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。33>- query、key、value参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。
31>- Q\_S和S1表示query shape中的S,KV\_S和S2表示key shape中的S,Q\_N和N1表示num\_query\_heads,KV\_N和N2表示num\_key\_value\_heads,T1表示query shape中的T,T2表示key shape中的输入样本序列长度的累加和。34>- Q\_S和S1表示query shape中的S,KV\_S和S2表示key shape中的S,Q\_N和N1表示num\_query\_heads,KV\_N和N2表示num\_key\_value\_heads,T1表示query shape中的T,T2表示key shape中的输入样本序列长度的累加和。
32-- **query**(`Tensor`):必选参数,对应公式中的$Q$,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。`layout_query`为BSND时shape为[B,S1,N1,D],当`layout_query`为TND时shape为[T1,N1,D],其中N1支持1/2/4/8/16/32/64/128。35+>
33-- **key**(`Tensor`):必选参数,对应公式中的$\tilde{K}$,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`,`layout_kv`时shape为[block\_num, block\_size, KV\_N, D],其中block\_num为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的倍数,最大支持1024。`layout_kv`为BSND时shape为[B, S2, KV\_N, D],`layout_kv`为TND时shape为[T2, KV\_N, D],其中KV\_N只支持1。36+- **query**(`Tensor`):必选参数,对应公式中的$Q$,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。`layout_query`为BSND时shape为[B,S1,N1,D],`layout_query`为TND时shape为[T1,N1,D],其中N1支持1/2/4/8/16/32/64/128
37+- **key**`Tensor`):必选参数,对应公式中的$\tilde{K}$,不支持非连续,数据格式支持ND,数据类型支持`bfloat16``float16``layout_kv`时shape为[block\_num, block\_size, KV\_N, D],其中block\_num为PageAttention时block总数,block\_size为一个block的token数,block\_size取值为16的倍数,最大支持1024。`layout_kv`为BSND时shape为[B, S2, KV\_N, D],`layout_kv`为TND时shape为[T2, KV\_N, D],其中KV\_N只支持1。
34 38 
35-- **value**(`Tensor`):必选参数,不支持非连续,对应公式中的$\tilde{V}$,维度N只支持1,数据格式支持ND,数据类型支持`bfloat16`和`float16`,shape与`key`的shape一致。39+- **value**(`Tensor`):必选参数,不支持非连续,对应公式中的$\tilde{V}$,维度N只支持1,数据格式支持ND,数据类型支持`bfloat16`和`float16`,shape与`key`的shape一致。
36 40
37-- **sparse\_indices**(`Tensor`):必选参数,代表离散取kvCache的索引,不支持非连续,数据格式支持ND,数据类型支持`int32`。当`layout_query`为BSND时,shape需要传入[B, Q\_S, KV\_N, sparse\_size],当`layout_query`为TND时,shape需要传入[Q\_T, KV\_N, sparse\_size],其中sparse\_size为一次离散选取的block数,需要保证每行有效值均在前半部分,无效值均在后半部分,且需要满足sparse\_size大于0。41+- **sparse\_indices**(`Tensor`):必选参数,代表离散取kvCache的索引,不支持非连续,数据格式支持ND,数据类型支持`int32`。当`layout_query`为BSND时,shape需要传入[B, Q\_S, KV\_N, sparse\_size],当`layout_query`为TND时,shape需要传入[Q\_T, KV\_N, sparse\_size],其中sparse\_size为一次离散选取的block数,需要保证每行有效值均在前半部分,无效值均在后半部分,且需要满足sparse\_size大于0。
38 42 
39-- **scale\_value**(`double`):必选参数,代表缩放系数,作为query和key矩阵乘后Muls的scalar值,数据类型支持`double`。43+- **scale\_value**(`double`):必选参数,代表缩放系数,作为query和key矩阵乘后Muls的scalar值,数据类型支持`double`。
40 44 
41- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。45- <strong>*</strong>:必选参数,代表其之前的变量是位置相关的,必须按照顺序输入;之后的变量是可选参数,位置无关,需要使用键值对赋值,不赋值会使用默认值。
42 46 
43-- **block\_table**(`Tensor`):可选参数,表示PageAttention中kvCache存储使用的block映射表。数据格式支持ND,数据类型支持`int32`,shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的S2对应的block数量,即S2\_max / block\_size向上取整。47+- **block\_table**(`Tensor`):可选参数,表示PageAttention中kvCache存储使用的block映射表。数据格式支持ND,数据类型支持`int32`,shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的S2对应的block数量,即S2\_max / block\_size向上取整。
44 48 
45-- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。该入参中每个Batch的有效token数不超过`query`中的维度S大小且不小于0。支持长度为B的一维tensor。<br>当`layout_query`为TND时,该入参必须传入,且以该入参元素的数量作为B值,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。49+- **actual\_seq\_lengths\_query**(`Tensor`):可选参数,表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。该入参中每个Batch的有效token数不超过`query`中的维度S大小且不小于0。支持长度为B的一维tensor。<br>当`layout_query`为TND时,该入参必须传入,且以该入参元素的数量作为B值,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。
46 50 
47-- **actual\_seq\_lengths\_kv**(`Tensor`):可选参数,表示不同Batch中`key`和`value`的有效token数,数据类型支持`int32`。如果不指定None,表示和key的shape的S长度相同。该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。<br>当`layout_kv`为TND或PA_BSND时,该入参必须传入,`layout_kv`为TND,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。51+- **actual\_seq\_lengths\_kv**(`Tensor`):可选参数,表示不同Batch中`key`和`value`的有效token数,数据类型支持`int32`。如果不指定None,表示和key的shape的S长度相同。该参数中每个Batch的有效token数不超过`key/value`中的维度S大小且不小于0。支持长度为B的一维tensor。<br>当`layout_kv`为TND或PA_BSND时,该入参必须传入,`layout_kv`为TND,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。
48 52 
49-- **query\_rope**(`Tensor`):可选参数,表示MLA结构中的query的rope信息,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。53+- **query\_rope**(`Tensor`):可选参数,表示MLA结构中的query的rope信息,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。
50 54
51-- **key\_rope**(`Tensor`):可选参数,表示MLA结构中的key的rope信息,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。55+- **key\_rope**(`Tensor`):可选参数,表示MLA结构中的key的rope信息,不支持非连续,数据格式支持ND,数据类型支持`bfloat16`和`float16`。
52 56 
53-- **sparse\_block\_size**(`int`):可选参数,代表sparse阶段的block大小,在计算importance score时使用,数据类型支持`int64`,取值范围为[1,128],且为2的幂次方。57+- **sparse\_block\_size**(`int`):可选参数,代表sparse阶段的block大小,在计算importance score时使用,数据类型支持`int64`,取值范围为[1,128],且为2的幂次方。
54- - sparse_block_size为1时,为Token-wise稀疏化场景,将每个token视为独立单元,在计算重要性分数时,评估每个查询token与每个键值token之间的独立关联程度。58+ - sparse_block_size为1时,为Token-wise稀疏化场景,将每个token视为独立单元,在计算重要性分数时,评估每个查询token与每个键值token之间的独立关联程度。
55- - sparse_block_size为大于1小于等于128时,为Block-wise稀疏化场景,将token序列划分为固定大小的连续块,以块为单位进行重要性评估,块内token共享相同的稀疏化决策。59+ - sparse_block_size为大于1小于等于128时,为Block-wise稀疏化场景,将token序列划分为固定大小的连续块,以块为单位进行重要性评估,块内token共享相同的稀疏化决策。
56 60 
57-- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,用户不特意指定时可传入默认值"BSND",支持传入BSND和TND。61+- **layout\_query**(`str`):可选参数,用于标识输入`query`的数据排布格式,用户不特意指定时可传入默认值"BSND",支持传入BSND和TND。
58 62 
59-- **layout\_kv**(`str`):可选参数,用于标识输入`key`的数据排布格式,用户不特意指定时可传入默认值"BSND",支持传入TND、BSND和PA\_BSND,其中PA\_BSND在使能PageAttention时使用。63+- **layout\_kv**(`str`):可选参数,用于标识输入`key`的数据排布格式,用户不特意指定时可传入默认值"BSND",支持传入TND、BSND和PA\_BSND,其中PA\_BSND在使能PageAttention时使用。
60 64 
61-- **sparse\_mode**(`int`):可选参数,表示sparse的模式。数据类型支持`int64`。65+- **sparse\_mode**(`int`):可选参数,表示sparse的模式。数据类型支持`int64`。
62- - sparse\_mode为0时,代表全部计算。66+ - sparse\_mode为0时,代表全部计算。
63- - sparse\_mode为3时,代表rightDownCausal模式的mask,对应以右下顶点往左上为划分线的下三角场景。67+ - sparse\_mode为3时,代表rightDownCausal模式的mask,对应以右下顶点往左上为划分线的下三角场景。
64 68 
65-- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。69+- **pre\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和前几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。
66 70 
67-- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。71+- **next\_tokens**(`int`):可选参数,用于稀疏计算,表示attention需要和后几个Token计算关联。数据类型支持`int64`,仅支持默认值2^63-1。
68 72 
69-- **attention\_mode**(`int`):可选参数,表示attention的模式,数据类型支持`int64`,仅支持传入2,表示MLA-absorb模式,即计算过程中会将query和key的nope部分分别和query_rope和key_rope的rope部分沿头维度(D)拼接,合并形成最终的query和key用于后续计算,且key和value共享同一份底层张量数据。73+- **attention\_mode**(`int`):可选参数,表示attention的模式,数据类型支持`int64`,仅支持传入2,表示MLA-absorb模式,即计算过程中会将query和key的nope部分分别和query_rope和key_rope的rope部分沿头维度(D)拼接,合并形成最终的query和key用于后续计算,且key和value共享同一份底层张量数据。
70 74 
71-- **return\_softmax\_lse**(`bool`):可选参数,用于表示是否返回softmax_max和softmax_sum。True表示返回,但图模式下不支持,False表示不返回;默认值为False。该参数仅在训练且`layout_kv`不为PA_BSND场景支持。75+- **return\_softmax\_lse**(`bool`):可选参数,用于表示是否返回softmax_max和softmax_sum。True表示返回,但图模式下不支持,False表示不返回;默认值为False。该参数仅在训练且`layout_kv`不为PA_BSND场景支持。
72 76 
73## 返回值说明77## 返回值说明
74 78 
75-- **attention\_out**(`Tensor`):公式中的输出。数据格式支持ND,数据类型支持`bfloat16`和`float16`。当layout\_query为BSND时shape为[B,S1,N1,D],当layout\_query为TND时shape为[T1,N1,D]。79+- **attention\_out**(`Tensor`):公式中的输出。数据格式支持ND,数据类型支持`bfloat16`和`float16`。当layout\_query为BSND时shape为[B,S1,N1,D],当layout\_query为TND时shape为[T1,N1,D]。
76-- **softmax\_max**(`Tensor`):可选输出,Attention算法对query乘key的结果,取max得到softmax_max,数据类型支持`float`。当layout\_query为BSND时shape为[B,N2,S1,N1/N2],当layout\_query为TND时shape为[N2,T1,N1/N2]。80+- **softmax\_max**(`Tensor`):可选输出,Attention算法对query乘key的结果,取max得到softmax_max,数据类型支持`float`。当layout\_query为BSND时shape为[B,N2,S1,N1/N2],当layout\_query为TND时shape为[N2,T1,N1/N2]。
77-- **softmax\_sum**(`Tensor`):可选输出,Attention算法query乘key的结果减去softmax_max, 再取exp,接着求sum,得到softmax_sum,数据类型支持`float`。当layout\_query为BSND时shape为[B,N2,S1,N1/N2],当layout\_query为TND时shape为[N2,T1,N1/N2]。81+- **softmax\_sum**(`Tensor`):可选输出,Attention算法query乘key的结果减去softmax_max, 再取exp,接着求sum,得到softmax_sum,数据类型支持`float`。当layout\_query为BSND时shape为[B,N2,S1,N1/N2],当layout\_query为TND时shape为[N2,T1,N1/N2]。
78 82 
79## 约束说明83## 约束说明
80 84 
81-- 该接口支持推理场景下使用。85+- 该接口支持推理场景下使用。
82-- 该接口支持图模式。86+- 该接口支持图模式。
83-- 参数query中的D和key、value的D值相等为512,参数query\_rope中的D和key\_rope的D值相等为64。87+- 参数query中的D和key、value的D值相等为512,参数query\_rope中的D和key\_rope的D值相等为64。
84-- 参数query、key、value的数据类型必须保持一致。88+- 参数query、key、value的数据类型必须保持一致。
85-- 支持sparse\_block\_size整除block\_size。89+- 支持sparse\_block\_size整除block\_size。
86-- `layout_kv`为PA_BSND时,`layout_query`和`layout_kv`无需一致; `layout_kv`为BSND或TND时,`layout_query`和`layout_kv`需保持一致。90+- `layout_kv`为PA_BSND时,`layout_query`和`layout_kv`无需一致; `layout_kv`为BSND或TND时,`layout_query`和`layout_kv`需保持一致。
87 91 
88## 调用示例92## 调用示例
93+ 
89- 单算子模式调用94- 单算子模式调用
90 95 
91 ```python96 ```python
@@ -139,6 +144,7 @@ torch_npu.npu_sparse_flash_attention(query, key, value, sparse_indices, scale_va
139 layout_query='BSND', layout_kv='BSND', sparse_mode=3, pre_tokens=(1<<63)-1, next_tokens=(1<<63)-1,144 layout_query='BSND', layout_kv='BSND', sparse_mode=3, pre_tokens=(1<<63)-1, next_tokens=(1<<63)-1,
140 attention_mode = attention_mode, return_softmax_lse = return_softmax_lse)145 attention_mode = attention_mode, return_softmax_lse = return_softmax_lse)
141 ```146 ```
147+ 
142- 图模式调用148- 图模式调用
143 149 
144 ```python150 ```python
@@ -217,4 +223,3 @@ torch_npu.npu_sparse_flash_attention(query, key, value, sparse_indices, scale_va
217 layout_query='BSND', layout_kv='BSND', sparse_mode=3, pre_tokens=(1<<63)-1, next_tokens=(1<<63)-1,223 layout_query='BSND', layout_kv='BSND', sparse_mode=3, pre_tokens=(1<<63)-1, next_tokens=(1<<63)-1,
218 attention_mode = attention_mode, return_softmax_lse = return_softmax_lse)224 attention_mode = attention_mode, return_softmax_lse = return_softmax_lse)
219 ```225 ```
220- 
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_sparse_lightning_indexer_grad_kl_loss.md+12-11
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_swiglu_quant.md+26-25
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_top_k_top_p.md+3-3
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_top_k_top_p_sample.md+32-26
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_trans_quant_param.md+2-1
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_transpose_batchmatmul.md+21-21
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_weight_quant_batchmatmul.md+63-63
Mdocs/zh/custom_APIs/torch_npu/torch_npu-reset_stream_limit.md+3-2
Mdocs/zh/custom_APIs/torch_npu/torch_npu-save_npugraph_tensor.md+1-1
Mdocs/zh/custom_APIs/torch_npu/torch_npu-scatter_update.md+3-2
Mdocs/zh/custom_APIs/torch_npu/torch_npu-scatter_update_.md+3-2
Mdocs/zh/custom_APIs/torch_npu/torch_npu-set_device_limit.md+3-3
Mdocs/zh/custom_APIs/torch_npu/torch_npu-set_stream_limit.md+13-5
Mdocs/zh/custom_APIs/torch_npu/torch_npu.md+0-1
Mdocs/zh/custom_APIs/torch_npu/torch_npu_list.md+1-5
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-_npu_dropout.md+3-4
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-copy_memory_.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-empty_with_format.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-fast_gelu.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_alloc_float_status.md+2-4
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_anchor_response_flags.md+3-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_apply_adam.md+2-1
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_batch_nms.md+15-19
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_bert_apply_adam.md+16-18
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_bmmV2.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_bounding_box_decode.md+3-5
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_bounding_box_encode.md+3-5
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_broadcast.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_ciou.md+3-6
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_clear_float_status.md+2-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_confusion_transpose.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_conv2d.md+1-3
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_conv3d.md+2-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_conv_transpose2d.md+2-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_convolution.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_convolution_transpose.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_deformable_conv2d.md+3-4
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_diou.md+3-6
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_dropout_with_add_softmax.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_dtype_cast.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_format_cast.md+2-1
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_format_cast_.md+3-1
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_fused_attention_score.md+2-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_get_float_status.md+2-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_giou.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_grid_assign_positive.md+3-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_gru.md+2-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_indexing.md+2-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_iou.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_layer_norm_eval.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_linear.md+1-3
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_lstm.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_max.md+2-6
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_min.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_mish.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_multi_head_attention.md+1-4
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_nms_rotated.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_nms_v4.md+2-3
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_nms_with_mask.md+7-8
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_one_hot.md+3-3
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_pad.md+3-5
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_ps_roi_pooling.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_ptiou.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_random_choice_with_mask.md+2-1
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_reshape.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_rms_norm.md+4-3
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_roi_align.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_rotated_iou.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_rotated_overlaps.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_sign_bits_pack.md+4-4
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_sign_bits_unpack.md+7-7
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_silu.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_slice.md+3-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_softmax_cross_entropy_with_logits.md+3-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_sort_v2.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_swiglu.md+4-4
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_transpose.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-npu_yolo_boxes_encode.md+1-2
Mdocs/zh/custom_APIs/torch_npu/(beta)torch_npu-one_.md+1-2
Mdocs/zh/install.md+27-19