已合并
fix(indexer_quant_cache): reject odd cache headDim for MX-FP4 mode #9613
wangxun21创建于 21 天前
fix(indexer_quant_cache): reject odd cache headDim for MX-FP4 mode #9613
已合并
共 4 个文件变更+324-291
| @@ -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。 |
| @@ -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 | -} | ||
| 39 | int64_t RoundUp(int64_t x, int64_t y) | 32 | int64_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 | - | ||
| 45 | constexpr int64_t INPUT_CACHE_IDX = 0; | 37 | constexpr int64_t INPUT_CACHE_IDX = 0; |
| 46 | constexpr int64_t INPUT_SCALE_IDX = 1; | 38 | constexpr int64_t INPUT_SCALE_IDX = 1; |
| 47 | constexpr int64_t INPUT_X_IDX = 2; | 39 | constexpr int64_t INPUT_X_IDX = 2; |
| @@ -54,7 +46,7 @@ constexpr int64_t INPUT_SLOT_MAPPING_IDX = 3; | |||
| 54 | constexpr size_t CACHE_VIEW_DIM_NUM = 4; | 46 | constexpr size_t CACHE_VIEW_DIM_NUM = 4; |
| 55 | constexpr size_t CACHE_BLOCKNUM_DIM = 0; | 47 | constexpr size_t CACHE_BLOCKNUM_DIM = 0; |
| 56 | constexpr size_t CACHE_BLOCKSIZE_DIM = 1; | 48 | constexpr size_t CACHE_BLOCKSIZE_DIM = 1; |
| 57 | -constexpr size_t CACHE_ONE_DIM = 2; // 倒数第二维, 必须 == 1 | 49 | +constexpr size_t CACHE_ONE_DIM = 2; // 倒数第二维, 必须 == 1 |
| 58 | constexpr int64_t CACHE_ONE_DIM_VALUE = 1; | 50 | constexpr int64_t CACHE_ONE_DIM_VALUE = 1; |
| 59 | constexpr int64_t ATTR_QUANT_MODE_INDEX = 0; | 51 | constexpr int64_t ATTR_QUANT_MODE_INDEX = 0; |
| 60 | constexpr int64_t ATTR_ROUND_SCALE_INDEX = 1; | 52 | constexpr int64_t ATTR_ROUND_SCALE_INDEX = 1; |
| @@ -63,7 +55,7 @@ constexpr int64_t BLOCK_SIZE = 32; | |||
| 63 | constexpr int64_t D_LENGTH_FULL_LOAD = 8192; | 55 | constexpr int64_t D_LENGTH_FULL_LOAD = 8192; |
| 64 | constexpr int64_t REPEAT_SIZE = 256; | 56 | constexpr int64_t REPEAT_SIZE = 256; |
| 65 | constexpr int64_t DOUBLE_BUFFER = 2; | 57 | constexpr int64_t DOUBLE_BUFFER = 2; |
| 66 | -constexpr int64_t FP4_PACK_NUM = 2; // MX-FP4: 2 fp4 values packed per byte | 58 | +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进行切分 |
| 68 | constexpr int64_t PER_BLOCK_FP16 = 128; | 60 | constexpr int64_t PER_BLOCK_FP16 = 128; |
| 69 | // MX-FP4 (quant_mode=3) 采用标准 MX 量化块, 每32个元素一个 e8m0 scale | 61 | // MX-FP4 (quant_mode=3) 采用标准 MX 量化块, 每32个元素一个 e8m0 scale |
| @@ -73,7 +65,7 @@ constexpr int64_t NORMAL_QUANT_MODE = 1; | |||
| 73 | constexpr int64_t HIFLOAT_QUANT_MODE = 2; | 65 | constexpr int64_t HIFLOAT_QUANT_MODE = 2; |
| 74 | constexpr int64_t MXFP4_QUANT_MODE = 3; | 66 | constexpr int64_t MXFP4_QUANT_MODE = 3; |
| 75 | constexpr int64_t SINGLE_ROW = 1; | 67 | constexpr int64_t SINGLE_ROW = 1; |
| 76 | -} | 68 | +} // namespace |
| 77 | 69 | ||
| 78 | ge::graphStatus IndexerQuantCacheTiling::GetPlatformInfo() | 70 | ge::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 | ||
| 98 | ge::graphStatus IndexerQuantCacheTiling::GetAttr() | 90 | ge::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 view | 141 | // 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 | - | ||
| 224 | ge::graphStatus IndexerQuantCacheTiling::CalcOpTiling() | 216 | ge::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) is | 317 | // 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() | |||
| 411 | ge::graphStatus IndexerQuantCacheTiling::PostTiling() | 418 | ge::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 | + | ||
| 421 | ge::graphStatus TilingPrepareForIndexerQuantCache(gert::TilingParseContext *context) | 439 | ge::graphStatus TilingPrepareForIndexerQuantCache(gert::TilingParseContext *context) |
| 422 | { | 440 | { |
| 423 | (void)context; | 441 | (void)context; |
| @@ -427,7 +445,7 @@ ge::graphStatus TilingPrepareForIndexerQuantCache(gert::TilingParseContext *cont | |||
| 427 | ge::graphStatus TilingForIndexerQuantCache(gert::TilingContext *context) | 445 | ge::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 optiling | 457 | +} // namespace optiling |
| @@ -16,7 +16,6 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | - | ||
| 20 | 19 | ||
| 21 | 20 | ||
| 22 | 21 | ||
| @@ -44,23 +43,23 @@ struct TilingOptionalParaInfo { | |||
| 44 | BEGIN_TILING_DATA_DEF(IndexerQuantCacheTilingData) | 43 | BEGIN_TILING_DATA_DEF(IndexerQuantCacheTilingData) |
| 45 | TILING_DATA_FIELD_DEF(int64_t, bs); | 44 | TILING_DATA_FIELD_DEF(int64_t, bs); |
| 46 | TILING_DATA_FIELD_DEF(int64_t, d); | 45 | TILING_DATA_FIELD_DEF(int64_t, d); |
| 47 | -TILING_DATA_FIELD_DEF(int64_t, scaleCol); // 一行多少个scale | 46 | +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处理行数 |
| 55 | TILING_DATA_FIELD_DEF(int64_t, quantMode); | 54 | TILING_DATA_FIELD_DEF(int64_t, quantMode); |
| 56 | TILING_DATA_FIELD_DEF(int64_t, roundScale); | 55 | TILING_DATA_FIELD_DEF(int64_t, roundScale); |
| 57 | TILING_DATA_FIELD_DEF(float, scalesAttr); | 56 | TILING_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 cache | 60 | +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 scale | 62 | +TILING_DATA_FIELD_DEF(int64_t, scaleBlockStride); // per-block (blockNum) stride of scale |
| 64 | END_TILING_DATA_DEF; | 63 | END_TILING_DATA_DEF; |
| 65 | 64 | ||
| 66 | REGISTER_TILING_DATA_CLASS(IndexerQuantCache, IndexerQuantCacheTilingData) | 65 | REGISTER_TILING_DATA_CLASS(IndexerQuantCache, IndexerQuantCacheTilingData) |
| @@ -74,7 +73,8 @@ struct IndexerQuantCacheCompileInfo { | |||
| 74 | // ----------算子Tiling入参信息解析及check类---------- | 73 | // ----------算子Tiling入参信息解析及check类---------- |
| 75 | class IndexerQuantCacheTiling { | 74 | class IndexerQuantCacheTiling { |
| 76 | public: | 75 | public: |
| 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 | + | ||
| 93 | private: | 96 | private: |
| 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 optiling | 125 | +} // namespace optiling |
| 123 | -#endif // INDEXER_QUANT_CACHE_TILING_ARCH35_H | 126 | +#endif // INDEXER_QUANT_CACHE_TILING_ARCH35_H |
| @@ -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::Test | 27 | +class IndexerQuantCacheTiling : public testing::Test { |
| 28 | -{ | ||
| 29 | protected: | 28 | protected: |
| 30 | static void SetUpTestCase() | 29 | static void SetUpTestCase() |
| 31 | { | 30 | { |
| @@ -43,23 +42,22 @@ protected: | |||
| 43 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_normal) | 42 | TEST_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) | |||
| 67 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_mxfp8) | 65 | TEST_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) | |||
| 91 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_mxfp4) | 88 | TEST_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, 其余保留)。 |
| 115 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_wide_headdim_ok) | 136 | TEST_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 | |||
| 140 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_hifloat_scale_non4d_rejected) | 160 | TEST_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 | |||
| 164 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_hifloat_scale_4d_ok) | 183 | TEST_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=1 | 188 | + {{{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) | |||
| 191 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_2d_rejected) | 209 | TEST_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) | |||
| 215 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_dim2_not_one_rejected) | 232 | TEST_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 | |||
| 239 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_cache_headdim_lt_d_rejected) | 255 | TEST_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 | |||
| 263 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_scale_headdim_lt_scalecol_rejected) | 278 | TEST_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 | |||
| 287 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_scale_non4d_rejected) | 301 | TEST_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 必须 4D | 306 | + {{{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) | |||
| 312 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_d_eq_8192_ok) | 325 | TEST_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) | |||
| 336 | TEST_F(IndexerQuantCacheTiling, indexer_quant_cache_tiling_d_gt_8192_rejected) | 348 | TEST_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 | } |