已合并
fix_issue_rms_normal_example #28
fix_issue_rms_normal_example #28
已合并
陈佳良创建于 3月14日
2 个文件变更+42-23
Mexamples/rms_norm/README.md+20-13
@@ -6,9 +6,8 @@
6- 调用方式:Kernel直调6- 调用方式:Kernel直调
7 7 
8 8 
9-## 样例支持AI处理器型号9+## 样例支持的产品
10-- Ascend 910C10+- 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 <= 716863 - 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```bash84```bash
78-cd ./examples85+output/bin/rms_norm --help // 查看帮助
79-bash run_examples.sh rms_norm86+output/bin/rms_norm --shape=16,32 // 运行样例
80```87```
Mexamples/rms_norm/rms_norm.cpp+22-10
@@ -23,7 +23,7 @@
23 23 
24static constexpr int32_t HEIGHT = 1;24static constexpr int32_t HEIGHT = 1;
25static constexpr int32_t WIDTH = 32;25static constexpr int32_t WIDTH = 32;
26-static constexpr int32_t MAX_DIM = 8;26+static constexpr int32_t MAX_DIM = 2;
27 27 
28template <typename T1, typename T2, typename T3>28template <typename T1, typename T2, typename T3>
29struct RmsNormConfig {29struct 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+ 
60struct Options {71struct 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 }