已合并
fix(indexer_quant_cache): reject odd cache headDim for MX-FP4 mode #9613
fix(indexer_quant_cache): reject odd cache headDim for MX-FP4 mode #9613
已合并
wangxun21创建于 21 天前
4 个文件变更+324-291
Mattention/indexer_quant_cache/docs/aclnnIndexerQuantCache.md+1-0
@@ -263,6 +263,7 @@ aclnnStatus aclnnIndexerQuantCache(
263- cache/cacheScale**仅在blockNum维支持非连续**(分页):各block可不紧密排布,但block内(blockSize、headDim维)须连续。263- cache/cacheScale**仅在blockNum维支持非连续**(分页):各block可不紧密排布,但block内(blockSize、headDim维)须连续。
264- **headDim长度约束**264- **headDim长度约束**
265 - **cache.headDim ≥ d**(MX-FP4模式以fp4元素计,d个fp4值占 ⌈d/2⌉ 字节)。265 - **cache.headDim ≥ d**(MX-FP4模式以fp4元素计,d个fp4值占 ⌈d/2⌉ 字节)。
266+ - **MX-FP4(quantMode=3)模式下,cache.headDim 必须为偶数**
266 - **cacheScale.headDim ≥ scaleCol**,scaleCol:MX-FP8/MX-FP4(quantMode=0/3)为 ⌈d/32⌉;Normal/HiFloat8(quantMode=1/2)为 1。267 - **cacheScale.headDim ≥ scaleCol**,scaleCol:MX-FP8/MX-FP4(quantMode=0/3)为 ⌈d/32⌉;Normal/HiFloat8(quantMode=1/2)为 1。
267 - 示例:d=128、quantMode=0 → scaleCol=4 → cache.headDim ≥ 128 且 cacheScale.headDim ≥ 4。268 - 示例:d=128、quantMode=0 → scaleCol=4 → cache.headDim ≥ 128 且 cacheScale.headDim ≥ 4。
268- x的最后一维(d轴)须能被32整除且 d ≤ 8192。269- x的最后一维(d轴)须能被32整除且 d ≤ 8192。
Mattention/indexer_quant_cache/op_host/indexer_quant_cache_tiling_arch35.cpp+69-51
@@ -29,19 +29,11 @@ int64_t CeilDiv(int64_t x, int64_t y)
29 }29 }
30 return x;30 return x;
31}31}
32-int64_t DownAlign(int64_t x, int64_t y)
33-{
34- if (y == 0) {
35- return x;
36- }
37- return (x / y) * y;
38-}
39int64_t RoundUp(int64_t x, int64_t y)32int64_t RoundUp(int64_t x, int64_t y)
40{33{
41 return CeilDiv(x, y) * y;34 return CeilDiv(x, y) * y;
42}35}
43 36 
44- 
45constexpr int64_t INPUT_CACHE_IDX = 0;37constexpr int64_t INPUT_CACHE_IDX = 0;
46constexpr int64_t INPUT_SCALE_IDX = 1;38constexpr int64_t INPUT_SCALE_IDX = 1;
47constexpr int64_t INPUT_X_IDX = 2;39constexpr int64_t INPUT_X_IDX = 2;
@@ -54,7 +46,7 @@ constexpr int64_t INPUT_SLOT_MAPPING_IDX = 3;
54constexpr size_t CACHE_VIEW_DIM_NUM = 4;46constexpr size_t CACHE_VIEW_DIM_NUM = 4;
55constexpr size_t CACHE_BLOCKNUM_DIM = 0;47constexpr size_t CACHE_BLOCKNUM_DIM = 0;
56constexpr size_t CACHE_BLOCKSIZE_DIM = 1;48constexpr size_t CACHE_BLOCKSIZE_DIM = 1;
57-constexpr size_t CACHE_ONE_DIM = 2; // 倒数第二维, 必须 == 149+constexpr size_t CACHE_ONE_DIM = 2; // 倒数第二维, 必须 == 1
58constexpr int64_t CACHE_ONE_DIM_VALUE = 1;50constexpr int64_t CACHE_ONE_DIM_VALUE = 1;
59constexpr int64_t ATTR_QUANT_MODE_INDEX = 0;51constexpr int64_t ATTR_QUANT_MODE_INDEX = 0;
60constexpr int64_t ATTR_ROUND_SCALE_INDEX = 1;52constexpr int64_t ATTR_ROUND_SCALE_INDEX = 1;
@@ -63,7 +55,7 @@ constexpr int64_t BLOCK_SIZE = 32;
63constexpr int64_t D_LENGTH_FULL_LOAD = 8192;55constexpr int64_t D_LENGTH_FULL_LOAD = 8192;
64constexpr int64_t REPEAT_SIZE = 256;56constexpr int64_t REPEAT_SIZE = 256;
65constexpr int64_t DOUBLE_BUFFER = 2;57constexpr int64_t DOUBLE_BUFFER = 2;
66-constexpr int64_t FP4_PACK_NUM = 2; // MX-FP4: 2 fp4 values packed per byte58+constexpr int64_t FP4_PACK_NUM = 2; // MX-FP4: 2 fp4 values packed per byte
67// per_block量化,每128个f16需要量化出一个scale, 因此切分尾轴时,以128为factor进行切分59// per_block量化,每128个f16需要量化出一个scale, 因此切分尾轴时,以128为factor进行切分
68constexpr int64_t PER_BLOCK_FP16 = 128;60constexpr int64_t PER_BLOCK_FP16 = 128;
69// MX-FP4 (quant_mode=3) 采用标准 MX 量化块, 每32个元素一个 e8m0 scale61// MX-FP4 (quant_mode=3) 采用标准 MX 量化块, 每32个元素一个 e8m0 scale
@@ -73,7 +65,7 @@ constexpr int64_t NORMAL_QUANT_MODE = 1;
73constexpr int64_t HIFLOAT_QUANT_MODE = 2;65constexpr int64_t HIFLOAT_QUANT_MODE = 2;
74constexpr int64_t MXFP4_QUANT_MODE = 3;66constexpr int64_t MXFP4_QUANT_MODE = 3;
75constexpr int64_t SINGLE_ROW = 1;67constexpr int64_t SINGLE_ROW = 1;
76-}68+} // namespace
77 69 
78ge::graphStatus IndexerQuantCacheTiling::GetPlatformInfo()70ge::graphStatus IndexerQuantCacheTiling::GetPlatformInfo()
79{71{
@@ -81,7 +73,7 @@ ge::graphStatus IndexerQuantCacheTiling::GetPlatformInfo()
81 if (platformInfo == nullptr) {73 if (platformInfo == nullptr) {
82 auto compileInfoPtr = context_->GetCompileInfo<IndexerQuantCacheCompileInfo>();74 auto compileInfoPtr = context_->GetCompileInfo<IndexerQuantCacheCompileInfo>();
83 OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_->GetNodeName(), "compile info is null"),75 OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_->GetNodeName(), "compile info is null"),
84- return ge::GRAPH_FAILED);76+ return ge::GRAPH_FAILED);
85 coreNum_ = compileInfoPtr->coreNum;77 coreNum_ = compileInfoPtr->coreNum;
86 ubSize_ = compileInfoPtr->ubSize;78 ubSize_ = compileInfoPtr->ubSize;
87 } else {79 } else {
@@ -97,7 +89,7 @@ ge::graphStatus IndexerQuantCacheTiling::GetPlatformInfo()
97 89 
98ge::graphStatus IndexerQuantCacheTiling::GetAttr()90ge::graphStatus IndexerQuantCacheTiling::GetAttr()
99{91{
100- auto* attrs = context_->GetAttrs();92+ auto *attrs = context_->GetAttrs();
101 OP_CHECK_NULL_WITH_CONTEXT(context_, attrs);93 OP_CHECK_NULL_WITH_CONTEXT(context_, attrs);
102 94 
103 auto quantMode = attrs->GetAttrPointer<int64_t>(ATTR_QUANT_MODE_INDEX);95 auto quantMode = attrs->GetAttrPointer<int64_t>(ATTR_QUANT_MODE_INDEX);
@@ -144,7 +136,7 @@ bool IndexerQuantCacheTiling::GetCacheViewLayout(
144 return true;136 return true;
145}137}
146 138 
147-ge::graphStatus IndexerQuantCacheTiling::ValidateCache4D(size_t inputIdx, const char *name, int64_t &lastDim)139+ge::graphStatus IndexerQuantCacheTiling::ValidateCache4D(size_t inputIdx, const char *inputName, int64_t &lastDim)
148{140{
149 // 4D-only 契约: 用逻辑(origin/view) shape 做维数门禁。连续 4D tensor 与 4D 分页 strided view141 // 4D-only 契约: 用逻辑(origin/view) shape 做维数门禁。连续 4D tensor 与 4D 分页 strided view
150 // 的 origin shape 均为该 4D 逻辑形状(GetShape()), 与 GetCacheViewLayout 经 inputShape->GetShape()142 // 的 origin shape 均为该 4D 逻辑形状(GetShape()), 与 GetCacheViewLayout 经 inputShape->GetShape()
@@ -153,15 +145,15 @@ ge::graphStatus IndexerQuantCacheTiling::ValidateCache4D(size_t inputIdx, const
153 OP_CHECK_NULL_WITH_CONTEXT(context_, shapePtr);145 OP_CHECK_NULL_WITH_CONTEXT(context_, shapePtr);
154 const auto &logical = shapePtr->GetShape();146 const auto &logical = shapePtr->GetShape();
155 OP_CHECK_IF(logical.GetDimNum() != CACHE_VIEW_DIM_NUM,147 OP_CHECK_IF(logical.GetDimNum() != CACHE_VIEW_DIM_NUM,
156- OP_LOGE(context_->GetNodeName(),148+ OP_LOGE(context_->GetNodeName(),
157- "%s must be 4D [blockNum, blockSize, 1, headDim], got dimNum=%zu",149+ "%s must be 4D [blockNum, blockSize, 1, headDim], got dimNum=%zu",
158- name, logical.GetDimNum()),150+ inputName, logical.GetDimNum()),
159- return ge::GRAPH_FAILED);151+ return ge::GRAPH_FAILED);
160 OP_CHECK_IF(logical.GetDim(CACHE_ONE_DIM) != CACHE_ONE_DIM_VALUE,152 OP_CHECK_IF(logical.GetDim(CACHE_ONE_DIM) != CACHE_ONE_DIM_VALUE,
161- OP_LOGE(context_->GetNodeName(),153+ OP_LOGE(context_->GetNodeName(),
162- "%s dim2 (second-to-last) must be 1 (one quantized vector per token), got %ld",154+ "%s dim2 (second-to-last) must be 1 (one quantized vector per token), got %ld",
163- name, logical.GetDim(CACHE_ONE_DIM)),155+ inputName, logical.GetDim(CACHE_ONE_DIM)),
164- return ge::GRAPH_FAILED);156+ return ge::GRAPH_FAILED);
165 // 连续场景的行宽(headDim)取末维; 分页 view 下随后由 view stride 覆盖, 此处仅作连续默认值。157 // 连续场景的行宽(headDim)取末维; 分页 view 下随后由 view stride 覆盖, 此处仅作连续默认值。
166 const auto &storage = shapePtr->GetStorageShape();158 const auto &storage = shapePtr->GetStorageShape();
167 lastDim = storage.GetDim(storage.GetDimNum() - 1);159 lastDim = storage.GetDim(storage.GetDimNum() - 1);
@@ -179,12 +171,13 @@ ge::graphStatus IndexerQuantCacheTiling::GetShapeAttrsInfoInner()
179 uint32_t xDims = xStorageShape.GetDimNum();171 uint32_t xDims = xStorageShape.GetDimNum();
180 uint32_t slotMappingDims = slotMappingShape.GetDimNum();172 uint32_t slotMappingDims = slotMappingShape.GetDimNum();
181 OP_CHECK_IF(xDims - 1 != slotMappingDims,173 OP_CHECK_IF(xDims - 1 != slotMappingDims,
182- OP_LOGE(context_->GetNodeName(), "slotMappingDims should equal xDims - 1"), return ge::GRAPH_FAILED);174+ OP_LOGE(context_->GetNodeName(), "slotMappingDims should equal xDims - 1"), return ge::GRAPH_FAILED);
183 int64_t bs = 1;175 int64_t bs = 1;
184 for (uint32_t i = 0; i < slotMappingDims; i++) {176 for (uint32_t i = 0; i < slotMappingDims; i++) {
185 int64_t temp = xStorageShape.GetDim(i);177 int64_t temp = xStorageShape.GetDim(i);
186- OP_CHECK_IF(temp != slotMappingShape.GetDim(i), OP_LOGE(context_->GetNodeName(),178+ OP_CHECK_IF(temp != slotMappingShape.GetDim(i),
187- "slotMappingShape should equal xStorageShape in dim %d", i), return ge::GRAPH_FAILED);179+ OP_LOGE(context_->GetNodeName(), "slotMappingShape should equal xStorageShape in dim %u", i),
180+ return ge::GRAPH_FAILED);
188 bs *= temp;181 bs *= temp;
189 }182 }
190 d_ = xStorageShape.GetDim(xDims - 1);183 d_ = xStorageShape.GetDim(xDims - 1);
@@ -193,20 +186,20 @@ ge::graphStatus IndexerQuantCacheTiling::GetShapeAttrsInfoInner()
193 // x 尾轴 d 必须 32 对齐: MX 模式每 32 个元素一个 scale, 且各量化分支按 32 元素块对齐处理,186 // x 尾轴 d 必须 32 对齐: MX 模式每 32 个元素一个 scale, 且各量化分支按 32 元素块对齐处理,
194 // 非 32 对齐会导致尾块读越界/scale 错位。187 // 非 32 对齐会导致尾块读越界/scale 错位。
195 OP_CHECK_IF(d_ <= 0 || d_ % BLOCK_SIZE != 0,188 OP_CHECK_IF(d_ <= 0 || d_ % BLOCK_SIZE != 0,
196- OP_LOGE(context_->GetNodeName(),189+ OP_LOGE(context_->GetNodeName(),
197- "the last dim (d) of x should be 32-aligned, got %ld", d_),190+ "the last dim (d) of x should be 32-aligned, got %ld", d_),
198- return ge::GRAPH_FAILED);191+ return ge::GRAPH_FAILED);
199 192 
200 OP_CHECK_IF(d_ > D_LENGTH_FULL_LOAD,193 OP_CHECK_IF(d_ > D_LENGTH_FULL_LOAD,
201- OP_LOGE(context_->GetNodeName(), "input x tail dimension must less than 8192, got %ld", d_),194+ OP_LOGE(context_->GetNodeName(), "input x tail dimension must less than 8192, got %ld", d_),
202- return ge::GRAPH_FAILED);195+ return ge::GRAPH_FAILED);
203 196 
204 OP_CHECK_IF(GetAttr() != ge::GRAPH_SUCCESS,197 OP_CHECK_IF(GetAttr() != ge::GRAPH_SUCCESS,
205- OP_LOGE(context_->GetNodeName(), "get attr failed."),198+ OP_LOGE(context_->GetNodeName(), "get attr failed."),
206- return ge::GRAPH_FAILED);199+ return ge::GRAPH_FAILED);
207 OP_CHECK_IF(quantMode_ < 0 || quantMode_ > MXFP4_QUANT_MODE,200 OP_CHECK_IF(quantMode_ < 0 || quantMode_ > MXFP4_QUANT_MODE,
208- OP_LOGE(context_->GetNodeName(), "quant_mode should be in [0,3], got %ld", quantMode_),201+ OP_LOGE(context_->GetNodeName(), "quant_mode should be in [0,3], got %ld", quantMode_),
209- return ge::GRAPH_FAILED);202+ return ge::GRAPH_FAILED);
210 203 
211 // 每行 scale 个数:204 // 每行 scale 个数:
212 // mode0 MX-FP8 / mode3 MX-FP4 : 每 32 个元素一个 scale (标准 MX 块)205 // mode0 MX-FP8 / mode3 MX-FP4 : 每 32 个元素一个 scale (标准 MX 块)
@@ -220,7 +213,6 @@ ge::graphStatus IndexerQuantCacheTiling::GetShapeAttrsInfoInner()
220 return ge::GRAPH_SUCCESS;213 return ge::GRAPH_SUCCESS;
221}214}
222 215 
223- 
224ge::graphStatus IndexerQuantCacheTiling::CalcOpTiling()216ge::graphStatus IndexerQuantCacheTiling::CalcOpTiling()
225{217{
226 rowOfFormerBlock_ = CeilDiv(bs_, static_cast<int64_t>(coreNum_));218 rowOfFormerBlock_ = CeilDiv(bs_, static_cast<int64_t>(coreNum_));
@@ -233,7 +225,7 @@ ge::graphStatus IndexerQuantCacheTiling::CalcOpTiling()
233 rowFactor_ = rowOnceLoop;225 rowFactor_ = rowOnceLoop;
234 int64_t scaleByteSize = 4;226 int64_t scaleByteSize = 4;
235 if (quantMode_ == MXFP8_QUANT_MODE || quantMode_ == MXFP4_QUANT_MODE) {227 if (quantMode_ == MXFP8_QUANT_MODE || quantMode_ == MXFP4_QUANT_MODE) {
236- scaleByteSize = 1; // MX 模式 scale 为 e8m0, 占1字节228+ scaleByteSize = 1; // MX 模式 scale 为 e8m0, 占1字节
237 }229 }
238 int64_t perBlockScaleElemNum = BLOCK_SIZE / scaleByteSize;230 int64_t perBlockScaleElemNum = BLOCK_SIZE / scaleByteSize;
239 int64_t xAlign = (quantMode_ == MXFP8_QUANT_MODE) ? PER_BLOCK_FP16 : 16;231 int64_t xAlign = (quantMode_ == MXFP8_QUANT_MODE) ? PER_BLOCK_FP16 : 16;
@@ -246,9 +238,10 @@ ge::graphStatus IndexerQuantCacheTiling::CalcOpTiling()
246 int64_t ySize = mid * RoundUp(d_, BLOCK_SIZE) * 1 * DOUBLE_BUFFER;238 int64_t ySize = mid * RoundUp(d_, BLOCK_SIZE) * 1 * DOUBLE_BUFFER;
247 int64_t scaleSize = mid * RoundUp(scaleCol_, perBlockScaleElemNum) * scaleByteSize * DOUBLE_BUFFER;239 int64_t scaleSize = mid * RoundUp(scaleCol_, perBlockScaleElemNum) * scaleByteSize * DOUBLE_BUFFER;
248 int64_t tmpBufferSize = RoundUp(mid, 8) * 4;240 int64_t tmpBufferSize = RoundUp(mid, 8) * 4;
249- int64_t mxScratchSize = (quantMode_ == MXFP8_QUANT_MODE)241+ int64_t mxScratchSize = 0;
250- ? (mid * RoundUp(d_, 8) * 4 + mid * CeilDiv(d_, 128) * 16 * 4)242+ if (quantMode_ == MXFP8_QUANT_MODE) {
251- : 0;243+ mxScratchSize = mid * RoundUp(d_, 8) * 4 + mid * CeilDiv(d_, 128) * 16 * 4;
244+ }
252 int64_t totalSize = xSize + ySize + scaleSize + tmpBufferSize + mxScratchSize;245 int64_t totalSize = xSize + ySize + scaleSize + tmpBufferSize + mxScratchSize;
253 if (totalSize <= static_cast<int64_t>(ubSize_)) {246 if (totalSize <= static_cast<int64_t>(ubSize_)) {
254 rowFactor_ = mid;247 rowFactor_ = mid;
@@ -280,6 +273,13 @@ ge::graphStatus IndexerQuantCacheTiling::CalcOpTiling()
280 if (ValidateCache4D(INPUT_CACHE_IDX, "cache", cacheLastDim) != ge::GRAPH_SUCCESS) {273 if (ValidateCache4D(INPUT_CACHE_IDX, "cache", cacheLastDim) != ge::GRAPH_SUCCESS) {
281 return ge::GRAPH_FAILED;274 return ge::GRAPH_FAILED;
282 }275 }
276+ // MX-FP4 每字节打包 2 个 fp4 元素, headDim 为奇数时行首会落在字节的高/低半字节,
277+ // 现有 kernel 仅支持整字节寻址(GlobalTensor<int8_t> + 整字节 CopyOut), 无法表达半字节偏移,
278+ // 因此在此处直接拒绝, 而非产出错误数据(与 ScatterPaKvCache 对 fp4 head_size 的校验一致)。
279+ if (CheckMxfp4EvenStride(cacheLastDim, "headDim") != ge::GRAPH_SUCCESS) {
280+ return ge::GRAPH_FAILED;
281+ }
282+ 
283 // 转为 kernel 寻址单位: MX-FP4 字节寻址打包 cache (÷2); 其余模式元素==字节, 直接取末维。283 // 转为 kernel 寻址单位: MX-FP4 字节寻址打包 cache (÷2); 其余模式元素==字节, 直接取末维。
284 cacheRowStride_ = (quantMode_ == MXFP4_QUANT_MODE) ? (cacheLastDim / FP4_PACK_NUM) : cacheLastDim;284 cacheRowStride_ = (quantMode_ == MXFP4_QUANT_MODE) ? (cacheLastDim / FP4_PACK_NUM) : cacheLastDim;
285 285 
@@ -292,7 +292,7 @@ ge::graphStatus IndexerQuantCacheTiling::CalcOpTiling()
292 return ge::GRAPH_FAILED;292 return ge::GRAPH_FAILED;
293 }293 }
294 scaleRowStride_ = scaleLastDim;294 scaleRowStride_ = scaleLastDim;
295- blockSize_ = 1; // default: contiguous flat slots (slot * rowStride)295+ blockSize_ = 1; // default: contiguous flat slots (slot * rowStride)
296 cacheBlockStride_ = cacheRowStride_;296 cacheBlockStride_ = cacheRowStride_;
297 scaleBlockStride_ = scaleRowStride_;297 scaleBlockStride_ = scaleRowStride_;
298 298 
@@ -317,6 +317,12 @@ ge::graphStatus IndexerQuantCacheTiling::CalcOpTiling()
317 // cache (int8 GM), so convert the cache strides to BYTES (÷2). scale (e8m0, 1 byte) is317 // cache (int8 GM), so convert the cache strides to BYTES (÷2). scale (e8m0, 1 byte) is
318 // already byte-equal so it is left unchanged.318 // already byte-equal so it is left unchanged.
319 if (quantMode_ == MXFP4_QUANT_MODE) {319 if (quantMode_ == MXFP4_QUANT_MODE) {
320+ // 分页 view 下同样以整字节寻址, row/block stride(fp4 元素)必须为偶数, 否则半字节偏移
321+ // 无法表达(与上面连续场景 cacheLastDim 的校验同理)。
322+ if (CheckMxfp4EvenStride(cacheRowStride_, "paged-view row stride") != ge::GRAPH_SUCCESS ||
323+ CheckMxfp4EvenStride(cacheBlockStride_, "paged-view block stride") != ge::GRAPH_SUCCESS) {
324+ return ge::GRAPH_FAILED;
325+ }
320 cacheRowStride_ = cacheRowStride_ / FP4_PACK_NUM;326 cacheRowStride_ = cacheRowStride_ / FP4_PACK_NUM;
321 cacheBlockStride_ = cacheBlockStride_ / FP4_PACK_NUM;327 cacheBlockStride_ = cacheBlockStride_ / FP4_PACK_NUM;
322 }328 }
@@ -326,16 +332,17 @@ ge::graphStatus IndexerQuantCacheTiling::CalcOpTiling()
326 // cache 行宽必须能容纳每行写出的 cacheCol, 否则散写越界到下一行。332 // cache 行宽必须能容纳每行写出的 cacheCol, 否则散写越界到下一行。
327 // (cacheCol 与 cacheRowStride_ 单位一致: MX-FP4 皆为字节, 其余皆为元素/字节)333 // (cacheCol 与 cacheRowStride_ 单位一致: MX-FP4 皆为字节, 其余皆为元素/字节)
328 OP_CHECK_IF(cacheRowStride_ < cacheCol,334 OP_CHECK_IF(cacheRowStride_ < cacheCol,
329- OP_LOGE(context_->GetNodeName(),335+ OP_LOGE(context_->GetNodeName(),
330- "cache headDim(row stride, in cache elements/bytes)=%ld must be >= per-token "336+ "cache headDim(row stride, in cache elements/bytes)=%ld must be >= per-token "
331- "quantized-x length=%ld", cacheRowStride_, cacheCol),337+ "quantized-x length=%ld",
332- return ge::GRAPH_FAILED);338+ cacheRowStride_, cacheCol),
339+ return ge::GRAPH_FAILED);
333 // cache_scale 行宽必须能容纳每行写出的 scaleCol(mode2: scaleCol=1), 否则散写越界到下一行。340 // cache_scale 行宽必须能容纳每行写出的 scaleCol(mode2: scaleCol=1), 否则散写越界到下一行。
334 OP_CHECK_IF(scaleRowStride_ < scaleCol_,341 OP_CHECK_IF(scaleRowStride_ < scaleCol_,
335- OP_LOGE(context_->GetNodeName(),342+ OP_LOGE(context_->GetNodeName(),
336- "cache_scale last dim(row stride)=%ld must be >= scaleCol=%ld",343+ "cache_scale last dim(row stride)=%ld must be >= scaleCol=%ld",
337- scaleRowStride_, scaleCol_),344+ scaleRowStride_, scaleCol_),
338- return ge::GRAPH_FAILED);345+ return ge::GRAPH_FAILED);
339 346 
340 tilingData_.set_bs(bs_);347 tilingData_.set_bs(bs_);
341 tilingData_.set_d(d_);348 tilingData_.set_d(d_);
@@ -411,13 +418,24 @@ ge::graphStatus IndexerQuantCacheTiling::GetWorkspaceSize()
411ge::graphStatus IndexerQuantCacheTiling::PostTiling()418ge::graphStatus IndexerQuantCacheTiling::PostTiling()
412{419{
413 context_->SetBlockDim(usedCoreNums_);420 context_->SetBlockDim(usedCoreNums_);
414- size_t* workspaces = context_->GetWorkspaceSizes(1);421+ size_t *workspaces = context_->GetWorkspaceSizes(1);
415 workspaces[0] = workspaceSize_;422 workspaces[0] = workspaceSize_;
416 tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());423 tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
417 context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());424 context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
418 return ge::GRAPH_SUCCESS;425 return ge::GRAPH_SUCCESS;
419}426}
420 427 
428+ge::graphStatus IndexerQuantCacheTiling::CheckMxfp4EvenStride(int64_t stride, const char *name)
429+{
430+ if (quantMode_ == MXFP4_QUANT_MODE && stride % FP4_PACK_NUM != 0) {
431+ const std::string strideStr = std::to_string(stride);
432+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(context_->GetNodeName(), name, strideStr.c_str(),
433+ "MX-FP4 cache stride in fp4 elements must be an even number.");
434+ return ge::GRAPH_FAILED;
435+ }
436+ return ge::GRAPH_SUCCESS;
437+}
438+ 
421ge::graphStatus TilingPrepareForIndexerQuantCache(gert::TilingParseContext *context)439ge::graphStatus TilingPrepareForIndexerQuantCache(gert::TilingParseContext *context)
422{440{
423 (void)context;441 (void)context;
@@ -427,7 +445,7 @@ ge::graphStatus TilingPrepareForIndexerQuantCache(gert::TilingParseContext *cont
427ge::graphStatus TilingForIndexerQuantCache(gert::TilingContext *context)445ge::graphStatus TilingForIndexerQuantCache(gert::TilingContext *context)
428{446{
429 OP_CHECK_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR("IndexerQuantCache", "Tiling context is null"),447 OP_CHECK_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR("IndexerQuantCache", "Tiling context is null"),
430- return ge::GRAPH_FAILED);448+ return ge::GRAPH_FAILED);
431 IndexerQuantCacheTiling IndexerQuantCacheTiling(context);449 IndexerQuantCacheTiling IndexerQuantCacheTiling(context);
432 return IndexerQuantCacheTiling.DoOpTiling();450 return IndexerQuantCacheTiling.DoOpTiling();
433}451}
@@ -436,4 +454,4 @@ IMPL_OP_OPTILING(IndexerQuantCache)
436 .Tiling(TilingForIndexerQuantCache)454 .Tiling(TilingForIndexerQuantCache)
437 .TilingParse<IndexerQuantCacheCompileInfo>(TilingPrepareForIndexerQuantCache);455 .TilingParse<IndexerQuantCacheCompileInfo>(TilingPrepareForIndexerQuantCache);
438 456 
439-} // namespace optiling457+} // namespace optiling
Mattention/indexer_quant_cache/op_host/indexer_quant_cache_tiling_arch35.h+20-17
@@ -16,7 +16,6 @@
16#ifndef INDEXER_QUANT_CACHE_TILING_ARCH35_H16#ifndef INDEXER_QUANT_CACHE_TILING_ARCH35_H
17#define INDEXER_QUANT_CACHE_TILING_ARCH35_H17#define INDEXER_QUANT_CACHE_TILING_ARCH35_H
18 18 
19- 
20#include <vector>19#include <vector>
21#include <iostream>20#include <iostream>
22#include "register/op_impl_registry.h"21#include "register/op_impl_registry.h"
@@ -44,23 +43,23 @@ struct TilingOptionalParaInfo {
44BEGIN_TILING_DATA_DEF(IndexerQuantCacheTilingData)43BEGIN_TILING_DATA_DEF(IndexerQuantCacheTilingData)
45TILING_DATA_FIELD_DEF(int64_t, bs);44TILING_DATA_FIELD_DEF(int64_t, bs);
46TILING_DATA_FIELD_DEF(int64_t, d);45TILING_DATA_FIELD_DEF(int64_t, d);
47-TILING_DATA_FIELD_DEF(int64_t, scaleCol); // 一行多少个scale46+TILING_DATA_FIELD_DEF(int64_t, scaleCol); // 一行多少个scale
48-TILING_DATA_FIELD_DEF(int64_t, rowOfFormerBlock); // 头核共需要处理多少行47+TILING_DATA_FIELD_DEF(int64_t, rowOfFormerBlock); // 头核共需要处理多少行
49-TILING_DATA_FIELD_DEF(int64_t, rowOfTailBlock); // 尾核共需要处理多少行48+TILING_DATA_FIELD_DEF(int64_t, rowOfTailBlock); // 尾核共需要处理多少行
50-TILING_DATA_FIELD_DEF(int64_t, rowLoopOfFormerBlock); // 头核需要几次ub搬入49+TILING_DATA_FIELD_DEF(int64_t, rowLoopOfFormerBlock); // 头核需要几次ub搬入
51-TILING_DATA_FIELD_DEF(int64_t, rowLoopOfTailBlock); // 尾核需要几次ub搬入50+TILING_DATA_FIELD_DEF(int64_t, rowLoopOfTailBlock); // 尾核需要几次ub搬入
52-TILING_DATA_FIELD_DEF(int64_t, rowFactor); // ub一次标准处理行数51+TILING_DATA_FIELD_DEF(int64_t, rowFactor); // ub一次标准处理行数
53-TILING_DATA_FIELD_DEF(int64_t, tailRowFactorOfFormerBlock); // 头核最后一次ub处理行数52+TILING_DATA_FIELD_DEF(int64_t, tailRowFactorOfFormerBlock); // 头核最后一次ub处理行数
54-TILING_DATA_FIELD_DEF(int64_t, tailRowFactorOfTailBlock); // 尾核最后一次ub处理行数53+TILING_DATA_FIELD_DEF(int64_t, tailRowFactorOfTailBlock); // 尾核最后一次ub处理行数
55TILING_DATA_FIELD_DEF(int64_t, quantMode);54TILING_DATA_FIELD_DEF(int64_t, quantMode);
56TILING_DATA_FIELD_DEF(int64_t, roundScale);55TILING_DATA_FIELD_DEF(int64_t, roundScale);
57TILING_DATA_FIELD_DEF(float, scalesAttr);56TILING_DATA_FIELD_DEF(float, scalesAttr);
58// 4D paged layout [blockNum, blockSize, 1, headDim]; blockNum non-contiguous.57// 4D paged layout [blockNum, blockSize, 1, headDim]; blockNum non-contiguous.
59-TILING_DATA_FIELD_DEF(int64_t, blockSize); // blockSize dim (1 => contiguous flat slots)58+TILING_DATA_FIELD_DEF(int64_t, blockSize); // blockSize dim (1 => contiguous flat slots)
60-TILING_DATA_FIELD_DEF(int64_t, cacheRowStride); // per-position stride of cache (= headDim elems)59+TILING_DATA_FIELD_DEF(int64_t, cacheRowStride); // per-position stride of cache (= headDim elems)
61-TILING_DATA_FIELD_DEF(int64_t, cacheBlockStride); // per-block (blockNum) stride of cache60+TILING_DATA_FIELD_DEF(int64_t, cacheBlockStride); // per-block (blockNum) stride of cache
62-TILING_DATA_FIELD_DEF(int64_t, scaleRowStride); // per-position stride of scale (= scaleCol)61+TILING_DATA_FIELD_DEF(int64_t, scaleRowStride); // per-position stride of scale (= scaleCol)
63-TILING_DATA_FIELD_DEF(int64_t, scaleBlockStride); // per-block (blockNum) stride of scale62+TILING_DATA_FIELD_DEF(int64_t, scaleBlockStride); // per-block (blockNum) stride of scale
64END_TILING_DATA_DEF;63END_TILING_DATA_DEF;
65 64 
66REGISTER_TILING_DATA_CLASS(IndexerQuantCache, IndexerQuantCacheTilingData)65REGISTER_TILING_DATA_CLASS(IndexerQuantCache, IndexerQuantCacheTilingData)
@@ -74,7 +73,8 @@ struct IndexerQuantCacheCompileInfo {
74// ----------算子Tiling入参信息解析及check类----------73// ----------算子Tiling入参信息解析及check类----------
75class IndexerQuantCacheTiling {74class IndexerQuantCacheTiling {
76public:75public:
77- explicit IndexerQuantCacheTiling(gert::TilingContext* tilingContext) : context_(tilingContext)76+ explicit IndexerQuantCacheTiling(gert::TilingContext *tilingContext)
77+ : context_(tilingContext)
78 {78 {
79 }79 }
80 ~IndexerQuantCacheTiling() = default;80 ~IndexerQuantCacheTiling() = default;
@@ -90,6 +90,9 @@ public:
90 // 4D-only 契约门禁: 校验 inputIdx 张量逻辑 shape 恰为 4D [blockNum, blockSize, 1, headDim]90 // 4D-only 契约门禁: 校验 inputIdx 张量逻辑 shape 恰为 4D [blockNum, blockSize, 1, headDim]
91 // (倒数第二维 == 1), 并回传其末维(headDim, 张量自身元素单位)。失败返回 GRAPH_FAILED。91 // (倒数第二维 == 1), 并回传其末维(headDim, 张量自身元素单位)。失败返回 GRAPH_FAILED。
92 ge::graphStatus ValidateCache4D(size_t inputIdx, const char *name, int64_t &lastDim);92 ge::graphStatus ValidateCache4D(size_t inputIdx, const char *name, int64_t &lastDim);
93+ // MX-FP4 模式下校验 stride(fp4元素单位) 必须为偶数, 否则半字节偏移无法表达
94+ ge::graphStatus CheckMxfp4EvenStride(int64_t stride, const char *name);
95+ 
93private:96private:
94 gert::TilingContext *context_ = nullptr;97 gert::TilingContext *context_ = nullptr;
95 IndexerQuantCacheTilingData tilingData_;98 IndexerQuantCacheTilingData tilingData_;
@@ -119,5 +122,5 @@ private:
119 int64_t tilingKey_ = 0;122 int64_t tilingKey_ = 0;
120};123};
121 124 
122-} // namespace optiling125+} // namespace optiling
123-#endif // INDEXER_QUANT_CACHE_TILING_ARCH35_H126+#endif // INDEXER_QUANT_CACHE_TILING_ARCH35_H
Mattention/indexer_quant_cache/tests/ut/op_host/test_indexer_quant_cache_tiling.cpp+234-223
@@ -24,8 +24,7 @@
24// 下列用例 x=[1024,128] (d=128): scaleCol mode0/3=4, mode1/2=1; cacheCol mode3=64, 其余=128。24// 下列用例 x=[1024,128] (d=128): scaleCol mode0/3=4, mode1/2=1; cacheCol mode3=64, 其余=128。
25// ---------------------------------------------------------------------------25// ---------------------------------------------------------------------------
26 26 
27-class IndexerQuantCacheTiling : public testing::Test27+class IndexerQuantCacheTiling : public testing::Test {
28-{
29protected:28protected:
30 static void SetUpTestCase()29 static void SetUpTestCase()
31 {30 {
@@ -43,23 +42,22 @@ protected:
43TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_normal)42TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_normal)
44{43{
45 gert::TilingContextPara optilingContextPara("IndexerQuantCache",44 gert::TilingContextPara optilingContextPara("IndexerQuantCache",
46- {45+ {
47- {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},46+ {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
48- {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},47+ {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},
49- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},48+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
50- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},49+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
51- },50+ },
52- {51+ {
53- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},52+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
54- {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},53+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
55- },54+ },
56- {55+ {
57- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},56+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},
58- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},57+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
59- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},58+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
60- },59+ },
61- nullptr, "Ascend950"60+ nullptr, "Ascend950");
62- );
63 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());61 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());
64}62}
65 63 
@@ -67,23 +65,22 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_normal)
67TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_mxfp8)65TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_mxfp8)
68{66{
69 gert::TilingContextPara optilingContextPara("IndexerQuantCache",67 gert::TilingContextPara optilingContextPara("IndexerQuantCache",
70- {68+ {
71- {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},69+ {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
72- {{{2048, 1, 1, 4}, {2048, 1, 1, 4}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},70+ {{{2048, 1, 1, 4}, {2048, 1, 1, 4}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
73- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},71+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
74- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},72+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
75- },73+ },
76- {74+ {
77- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},75+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
78- {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},76+ {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
79- },77+ },
80- {78+ {
81- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},79+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
82- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},80+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
83- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},81+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
84- },82+ },
85- nullptr, "Ascend950"83+ nullptr, "Ascend950");
86- );
87 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());84 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());
88}85}
89 86 
@@ -91,47 +88,70 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_mxfp8)
91TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_mxfp4)88TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_mxfp4)
92{89{
93 gert::TilingContextPara optilingContextPara("IndexerQuantCache",90 gert::TilingContextPara optilingContextPara("IndexerQuantCache",
94- {91+ {
95- {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND},92+ {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND},
96- {{{2048, 1, 1, 4}, {2048, 1, 1, 4}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},93+ {{{2048, 1, 1, 4}, {2048, 1, 1, 4}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
97- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},94+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
98- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},95+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
99- },96+ },
100- {97+ {
101- {{{}, {}}, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND},98+ {{{}, {}}, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND},
102- {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},99+ {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
103- },100+ },
104- {101+ {
105- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(3)},102+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(3)},
106- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},103+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
107- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},104+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
108- },105+ },
109- nullptr, "Ascend950"106+ nullptr, "Ascend950");
110- );
111 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());107 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());
112}108}
113 109 
110+// mode3 MX-FP4: cache headDim(fp4元素)为奇数(129) -> GRAPH_FAILED。
111+// FP4每字节打包2个元素, headDim为奇数时逐行起点会交替落在字节高/低半字节, 现有整字节寻址
112+// 实现无法表达该偏移, 故直接拒绝而非产出错误数据。
113+TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_mxfp4_odd_headdim_rejected)
114+{
115+ gert::TilingContextPara optilingContextPara("IndexerQuantCache",
116+ {
117+ {{{2048, 1, 1, 129}, {2048, 1, 1, 129}}, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND},
118+ {{{2048, 1, 1, 4}, {2048, 1, 1, 4}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
119+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
120+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
121+ },
122+ {
123+ {{{}, {}}, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND},
124+ {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
125+ },
126+ {
127+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(3)},
128+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
129+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
130+ },
131+ nullptr, "Ascend950");
132+ ExecuteTestCase(optilingContextPara, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());
133+}
134+ 
114// cache headDim > d (256 > 128) -> 成功 (Reading B: 行可更宽, 只写 d, 其余保留)。135// cache headDim > d (256 > 128) -> 成功 (Reading B: 行可更宽, 只写 d, 其余保留)。
115TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_wide_headdim_ok)136TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_wide_headdim_ok)
116{137{
117 gert::TilingContextPara optilingContextPara("IndexerQuantCache",138 gert::TilingContextPara optilingContextPara("IndexerQuantCache",
118- {139+ {
119- {{{2048, 1, 1, 256}, {2048, 1, 1, 256}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},140+ {{{2048, 1, 1, 256}, {2048, 1, 1, 256}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
120- {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},141+ {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},
121- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},142+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
122- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},143+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
123- },144+ },
124- {145+ {
125- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},146+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
126- {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},147+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
127- },148+ },
128- {149+ {
129- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},150+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},
130- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},151+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
131- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},152+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
132- },153+ },
133- nullptr, "Ascend950"154+ nullptr, "Ascend950");
134- );
135 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());155 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());
136}156}
137 157 
@@ -140,23 +160,22 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_wide_headdim_ok
140TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_hifloat_scale_non4d_rejected)160TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_hifloat_scale_non4d_rejected)
141{161{
142 gert::TilingContextPara optilingContextPara("IndexerQuantCache",162 gert::TilingContextPara optilingContextPara("IndexerQuantCache",
143- {163+ {
144- {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},164+ {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
145- {{{2048, 1}, {2048, 1}}, ge::DT_FLOAT, ge::FORMAT_ND}, // 2D scale, mode2 现已校验 -> 拒绝165+ {{{2048, 1}, {2048, 1}}, ge::DT_FLOAT, ge::FORMAT_ND}, // 2D scale, mode2 现已校验 -> 拒绝
146- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},166+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
147- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},167+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
148- },168+ },
149- {169+ {
150- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},170+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
151- {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},171+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
152- },172+ },
153- {173+ {
154- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(2)},174+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(2)},
155- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},175+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
156- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},176+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
157- },177+ },
158- nullptr, "Ascend950"178+ nullptr, "Ascend950");
159- );
160 ExecuteTestCase(optilingContextPara, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());179 ExecuteTestCase(optilingContextPara, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());
161}180}
162 181 
@@ -164,23 +183,22 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_hifloat_scale_non4d_r
164TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_hifloat_scale_4d_ok)183TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_hifloat_scale_4d_ok)
165{184{
166 gert::TilingContextPara optilingContextPara("IndexerQuantCache",185 gert::TilingContextPara optilingContextPara("IndexerQuantCache",
167- {186+ {
168- {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},187+ {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
169- {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND}, // 4D scale, scaleCol=1188+ {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND}, // 4D scale, scaleCol=1
170- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},189+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
171- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},190+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
172- },191+ },
173- {192+ {
174- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},193+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
175- {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},194+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
176- },195+ },
177- {196+ {
178- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(2)},197+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(2)},
179- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},198+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
180- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},199+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
181- },200+ },
182- nullptr, "Ascend950"201+ nullptr, "Ascend950");
183- );
184 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());202 ExecuteTestCase(optilingContextPara, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());
185}203}
186 204 
@@ -191,23 +209,22 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_hifloat_scale_4d_ok)
191TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_2d_rejected)209TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_2d_rejected)
192{210{
193 gert::TilingContextPara para("IndexerQuantCache",211 gert::TilingContextPara para("IndexerQuantCache",
194- {212+ {
195- {{{2048, 128}, {2048, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},213+ {{{2048, 128}, {2048, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
196- {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},214+ {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},
197- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},215+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
198- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},216+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
199- },217+ },
200- {218+ {
201- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},219+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
202- {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},220+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
203- },221+ },
204- {222+ {
205- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},223+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},
206- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},224+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
207- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},225+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
208- },226+ },
209- nullptr, "Ascend950"227+ nullptr, "Ascend950");
210- );
211 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());228 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());
212}229}
213 230 
@@ -215,23 +232,22 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_2d_rejected)
215TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_dim2_not_one_rejected)232TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_dim2_not_one_rejected)
216{233{
217 gert::TilingContextPara para("IndexerQuantCache",234 gert::TilingContextPara para("IndexerQuantCache",
218- {235+ {
219- {{{128, 16, 2, 128}, {128, 16, 2, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},236+ {{{128, 16, 2, 128}, {128, 16, 2, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
220- {{{128, 16, 1, 1}, {128, 16, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},237+ {{{128, 16, 1, 1}, {128, 16, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},
221- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},238+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
222- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},239+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
223- },240+ },
224- {241+ {
225- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},242+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
226- {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},243+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
227- },244+ },
228- {245+ {
229- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},246+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},
230- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},247+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
231- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},248+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
232- },249+ },
233- nullptr, "Ascend950"250+ nullptr, "Ascend950");
234- );
235 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());251 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());
236}252}
237 253 
@@ -239,23 +255,22 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_dim2_not_one_re
239TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_headdim_lt_d_rejected)255TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_headdim_lt_d_rejected)
240{256{
241 gert::TilingContextPara para("IndexerQuantCache",257 gert::TilingContextPara para("IndexerQuantCache",
242- {258+ {
243- {{{2048, 1, 1, 64}, {2048, 1, 1, 64}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},259+ {{{2048, 1, 1, 64}, {2048, 1, 1, 64}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
244- {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},260+ {{{2048, 1, 1, 1}, {2048, 1, 1, 1}}, ge::DT_FLOAT, ge::FORMAT_ND},
245- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},261+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
246- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},262+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
247- },263+ },
248- {264+ {
249- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},265+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
250- {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},266+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
251- },267+ },
252- {268+ {
253- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},269+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},
254- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},270+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
255- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},271+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
256- },272+ },
257- nullptr, "Ascend950"273+ nullptr, "Ascend950");
258- );
259 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());274 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());
260}275}
261 276 
@@ -263,23 +278,22 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_headdim_lt_d_re
263TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_scale_headdim_lt_scalecol_rejected)278TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_scale_headdim_lt_scalecol_rejected)
264{279{
265 gert::TilingContextPara para("IndexerQuantCache",280 gert::TilingContextPara para("IndexerQuantCache",
266- {281+ {
267- {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},282+ {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
268- {{{2048, 1, 1, 2}, {2048, 1, 1, 2}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND}, // 2 < scaleCol(4)283+ {{{2048, 1, 1, 2}, {2048, 1, 1, 2}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND}, // 2 < scaleCol(4)
269- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},284+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
270- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},285+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
271- },286+ },
272- {287+ {
273- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},288+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
274- {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},289+ {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
275- },290+ },
276- {291+ {
277- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},292+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
278- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},293+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
279- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},294+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
280- },295+ },
281- nullptr, "Ascend950"296+ nullptr, "Ascend950");
282- );
283 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());297 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());
284}298}
285 299 
@@ -287,23 +301,22 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_scale_headdim_lt_scal
287TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_scale_non4d_rejected)301TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_scale_non4d_rejected)
288{302{
289 gert::TilingContextPara para("IndexerQuantCache",303 gert::TilingContextPara para("IndexerQuantCache",
290- {304+ {
291- {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},305+ {{{2048, 1, 1, 128}, {2048, 1, 1, 128}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
292- {{{2048, 1}, {2048, 1}}, ge::DT_FLOAT, ge::FORMAT_ND}, // 2D scale, mode1 必须 4D306+ {{{2048, 1}, {2048, 1}}, ge::DT_FLOAT, ge::FORMAT_ND}, // 2D scale, mode1 必须 4D
293- {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},307+ {{{1024, 128}, {1024, 128}}, ge::DT_FLOAT16, ge::FORMAT_ND},
294- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},308+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
295- },309+ },
296- {310+ {
297- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},311+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
298- {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},312+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
299- },313+ },
300- {314+ {
301- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},315+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(1)},
302- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},316+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
303- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},317+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
304- },318+ },
305- nullptr, "Ascend950"319+ nullptr, "Ascend950");
306- );
307 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());320 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());
308}321}
309 322 
@@ -312,23 +325,22 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_scale_non4d_rejected)
312TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_d_eq_8192_ok)325TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_d_eq_8192_ok)
313{326{
314 gert::TilingContextPara para("IndexerQuantCache",327 gert::TilingContextPara para("IndexerQuantCache",
315- {328+ {
316- {{{2048, 1, 1, 8192}, {2048, 1, 1, 8192}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},329+ {{{2048, 1, 1, 8192}, {2048, 1, 1, 8192}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
317- {{{2048, 1, 1, 256}, {2048, 1, 1, 256}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},330+ {{{2048, 1, 1, 256}, {2048, 1, 1, 256}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
318- {{{1024, 8192}, {1024, 8192}}, ge::DT_FLOAT16, ge::FORMAT_ND},331+ {{{1024, 8192}, {1024, 8192}}, ge::DT_FLOAT16, ge::FORMAT_ND},
319- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},332+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
320- },333+ },
321- {334+ {
322- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},335+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
323- {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},336+ {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
324- },337+ },
325- {338+ {
326- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},339+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
327- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},340+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
328- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},341+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
329- },342+ },
330- nullptr, "Ascend950"343+ nullptr, "Ascend950");
331- );
332 ExecuteTestCase(para, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());344 ExecuteTestCase(para, ge::GRAPH_SUCCESS, std::numeric_limits<uint64_t>::max());
333}345}
334 346 
@@ -336,22 +348,21 @@ TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_d_eq_8192_ok)
336TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_d_gt_8192_rejected)348TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_d_gt_8192_rejected)
337{349{
338 gert::TilingContextPara para("IndexerQuantCache",350 gert::TilingContextPara para("IndexerQuantCache",
339- {351+ {
340- {{{2048, 1, 1, 8224}, {2048, 1, 1, 8224}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},352+ {{{2048, 1, 1, 8224}, {2048, 1, 1, 8224}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
341- {{{2048, 1, 1, 257}, {2048, 1, 1, 257}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},353+ {{{2048, 1, 1, 257}, {2048, 1, 1, 257}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
342- {{{1024, 8224}, {1024, 8224}}, ge::DT_FLOAT16, ge::FORMAT_ND},354+ {{{1024, 8224}, {1024, 8224}}, ge::DT_FLOAT16, ge::FORMAT_ND},
343- {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},355+ {{{1024}, {1024}}, ge::DT_INT32, ge::FORMAT_ND},
344- },356+ },
345- {357+ {
346- {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},358+ {{{}, {}}, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND},
347- {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},359+ {{{}, {}}, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND},
348- },360+ },
349- {361+ {
350- {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},362+ {"quant_mode", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
351- {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},363+ {"round_scale", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
352- {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},364+ {"x_scale", Ops::Transformer::AnyValue::CreateFrom<float>(1.0f)},
353- },365+ },
354- nullptr, "Ascend950"366+ nullptr, "Ascend950");
355- );
356 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());367 ExecuteTestCase(para, ge::GRAPH_FAILED, std::numeric_limits<uint64_t>::max());
357}368}