已合并
fix_issue_rms_normal_example #28
陈佳良创建于 3月14日
fix_issue_rms_normal_example #28
已合并
共 2 个文件变更+42-23
| @@ -6,9 +6,8 @@ | |||
| 6 | - 调用方式:Kernel直调 | 6 | - 调用方式:Kernel直调 |
| 7 | 7 | ||
| 8 | 8 | ||
| 9 | -## 样例支持AI处理器型号 | 9 | +## 样例支持的产品 |
| 10 | -- Ascend 910C | 10 | +- Ascend 950PR/Ascend 950DT |
| 11 | -- Ascend 910B | ||
| 12 | 11 | ||
| 13 | 12 | ||
| 14 | ## 算子描述 | 13 | ## 算子描述 |
| @@ -36,21 +35,21 @@ | |||
| 36 | </tr></thead> | 35 | </tr></thead> |
| 37 | <tbody> | 36 | <tbody> |
| 38 | <tr> | 37 | <tr> |
| 39 | - <td>x</td> | 38 | + <td>in1</td> |
| 40 | <td>输入</td> | 39 | <td>输入</td> |
| 41 | <td>表示进行归一化计算的输入。公式中的`x`。</td> | 40 | <td>表示进行归一化计算的输入。公式中的`x`。</td> |
| 42 | <td>float</td> | 41 | <td>float</td> |
| 43 | <td>ND</td> | 42 | <td>ND</td> |
| 44 | </tr> | 43 | </tr> |
| 45 | <tr> | 44 | <tr> |
| 46 | - <td>gamma</td> | 45 | + <td>in2</td> |
| 47 | <td>输入</td> | 46 | <td>输入</td> |
| 48 | <td>表示进行归一化计算的缩放因子(权重),公式中的`g`。</td> | 47 | <td>表示进行归一化计算的缩放因子(权重),公式中的`g`。</td> |
| 49 | <td>float</td> | 48 | <td>float</td> |
| 50 | <td>ND</td> | 49 | <td>ND</td> |
| 51 | </tr> | 50 | </tr> |
| 52 | <tr> | 51 | <tr> |
| 53 | - <td>y</td> | 52 | + <td>out</td> |
| 54 | <td>输出</td> | 53 | <td>输出</td> |
| 55 | <td>表示进行归一化后的最终输出,公式中的`RmsNorm(x)`。</td> | 54 | <td>表示进行归一化后的最终输出,公式中的`RmsNorm(x)`。</td> |
| 56 | <td>float</td> | 55 | <td>float</td> |
| @@ -59,7 +58,7 @@ | |||
| 59 | </tbody></table> | 58 | </tbody></table> |
| 60 | 规格说明: | 59 | 规格说明: |
| 61 | 60 | ||
| 62 | -- 当前只支持二维输入, | 61 | +- 当前只支持二维输入 |
| 63 | - 总的输入Shape(M, N)要满足: | 62 | - 总的输入Shape(M, N)要满足: |
| 64 | - M < 8160,N <= 7168 | 63 | - M < 8160,N <= 7168 |
| 65 | - N需要32元素对齐 | 64 | - N需要32元素对齐 |
| @@ -68,13 +67,21 @@ | |||
| 68 | 67 | ||
| 69 | ## 目录结构 | 68 | ## 目录结构 |
| 70 | 69 | ||
| 71 | -| 文件名 | 描述 | | 70 | +| 文件名 | 描述 | |
| 72 | -| ------------------------------------------------------------ | ------------------------------------------ | | 71 | +|------------------------|------------------| |
| 73 | -| [rms_norm.cpp](./rms_norm.cpp) | RmsNorm算子代码实现以及调用样例 | | 72 | +| [rms_norm.cpp](./rms_norm.cpp) | RmsNorm样例算子代码实现 | |
| 73 | +| [CMakeLists.txt](./CMakeLists.txt) | RmsNorm样例算子的编译构建文件 | | ||
| 74 | +| [README.md](./README.md) | RmsNorm样例算子的说明文档 | | ||
| 74 | 75 | ||
| 75 | -## 算子运行 | 76 | +## RmsNorm样例算子的编译和运行 |
| 77 | +- 编译 | ||
| 78 | +在代码仓根目录下执行: | ||
| 79 | +```bash | ||
| 80 | +bash scripts/build.sh -DSOC=ascend950 rms_norm | ||
| 81 | +``` | ||
| 82 | +- 运行 | ||
| 76 | 在代码仓目录下执行: | 83 | 在代码仓目录下执行: |
| 77 | ```bash | 84 | ```bash |
| 78 | -cd ./examples | 85 | +output/bin/rms_norm --help // 查看帮助 |
| 79 | -bash run_examples.sh rms_norm | 86 | +output/bin/rms_norm --shape=16,32 // 运行样例 |
| 80 | ``` | 87 | ``` |
| @@ -23,7 +23,7 @@ | |||
| 23 | 23 | ||
| 24 | static constexpr int32_t HEIGHT = 1; | 24 | static constexpr int32_t HEIGHT = 1; |
| 25 | static constexpr int32_t WIDTH = 32; | 25 | static constexpr int32_t WIDTH = 32; |
| 26 | -static constexpr int32_t MAX_DIM = 8; | 26 | +static constexpr int32_t MAX_DIM = 2; |
| 27 | 27 | ||
| 28 | template <typename T1, typename T2, typename T3> | 28 | template <typename T1, typename T2, typename T3> |
| 29 | struct RmsNormConfig { | 29 | struct RmsNormConfig { |
| @@ -46,17 +46,28 @@ struct RmsNormConfig { | |||
| 46 | } | 46 | } |
| 47 | }; | 47 | }; |
| 48 | 48 | ||
| 49 | - static constexpr Atvoss::Ele::DefaultBlockPolicy<TileShape> blockPolicy{TileShape{}}; | 49 | + static constexpr Atvoss::Ele::DefaultBlockPolicy<TileShape> blockPolicy { |
| 50 | - static constexpr Atvoss::Ele::DefaultKernelPolicy kernelPolicy{Atvoss::Ele::DefaultSegmentPolicy::UniformSegment}; | 50 | + TileShape{} |
| 51 | + }; | ||
| 52 | + static constexpr Atvoss::Ele::DefaultKernelPolicy kernelPolicy { | ||
| 53 | + Atvoss::Ele::DefaultSegmentPolicy::UniformSegment | ||
| 54 | + }; | ||
| 51 | 55 | ||
| 52 | using ArchTag = Atvoss::Arch::DAV_3510; | 56 | using ArchTag = Atvoss::Arch::DAV_3510; |
| 53 | - using BlockOp = Atvoss::Ele::BlockBuilder<RmsNormCompute, ArchTag, blockPolicy, Atvoss::Ele::DefaultBlockConfig>; | 57 | + using BlockOp = Atvoss::Ele::BlockBuilder< |
| 58 | + RmsNormCompute, | ||
| 59 | + ArchTag, | ||
| 60 | + blockPolicy, | ||
| 61 | + Atvoss::Ele::DefaultBlockConfig>; | ||
| 54 | 62 | ||
| 55 | - using KernelOp = Atvoss::Ele::KernelBuilder<BlockOp, kernelPolicy>; | 63 | + using KernelOp = Atvoss::Ele::KernelBuilder< |
| 64 | + BlockOp, | ||
| 65 | + kernelPolicy>; | ||
| 56 | 66 | ||
| 57 | using DeviceOp = Atvoss::DeviceAdapter<KernelOp>; | 67 | using DeviceOp = Atvoss::DeviceAdapter<KernelOp>; |
| 58 | }; | 68 | }; |
| 59 | 69 | ||
| 70 | + | ||
| 60 | struct Options { | 71 | struct Options { |
| 61 | // 存储解析后的值 | 72 | // 存储解析后的值 |
| 62 | std::vector<int> shape; | 73 | std::vector<int> shape; |
| @@ -64,7 +75,8 @@ struct Options { | |||
| 64 | 75 | ||
| 65 | // 默认构造:不解析,只设默认值 | 76 | // 默认构造:不解析,只设默认值 |
| 66 | Options() : shape({}), help(false) | 77 | Options() : shape({}), help(false) |
| 67 | - {} | 78 | + { |
| 79 | + } | ||
| 68 | 80 | ||
| 69 | // 解析函数 | 81 | // 解析函数 |
| 70 | void parse(int argc, char const* argv[]) | 82 | void parse(int argc, char const* argv[]) |
| @@ -87,10 +99,10 @@ struct Options { | |||
| 87 | << "\n" | 99 | << "\n" |
| 88 | << "Options:\n" | 100 | << "Options:\n" |
| 89 | << " --help Print this message\n" | 101 | << " --help Print this message\n" |
| 90 | - << " --shape=M,N,O,... Tensor dimensions (e.g., --shape=4,3,224,224)\n" | 102 | + << " --shape=M,N,... Tensor dimensions (e.g., --shape=512,32)\n" |
| 91 | << "\n" | 103 | << "\n" |
| 92 | << "Example:\n" | 104 | << "Example:\n" |
| 93 | - << " " << progName << " --shape=512,3 \n"; | 105 | + << " " << progName << " --shape=512,32 \n"; |
| 94 | } | 106 | } |
| 95 | 107 | ||
| 96 | // 打印当前配置 | 108 | // 打印当前配置 |
| @@ -116,7 +128,7 @@ private: | |||
| 116 | void validate(const char* progName) | 128 | void validate(const char* progName) |
| 117 | { | 129 | { |
| 118 | if (help) { | 130 | if (help) { |
| 119 | - return; // 帮助不需要验证 | 131 | + return; // 帮助不需要验证 |
| 120 | } | 132 | } |
| 121 | 133 | ||
| 122 | if (shape.empty()) { | 134 | if (shape.empty()) { |
| @@ -125,7 +137,7 @@ private: | |||
| 125 | exit(1); | 137 | exit(1); |
| 126 | } | 138 | } |
| 127 | if (shape.size() > MAX_DIM) { | 139 | if (shape.size() > MAX_DIM) { |
| 128 | - std::cerr << "[ERROR] Input shape max dim is 8, current shape dim is: " << shape.size() << "\n"; | 140 | + std::cerr << "[ERROR] Input shape max dim is 2, current shape dim is: " << shape.size() << "\n"; |
| 129 | PrintUsage(progName); | 141 | PrintUsage(progName); |
| 130 | exit(1); | 142 | exit(1); |
| 131 | } | 143 | } |