Thanks for sending an issue! Please fill in the following template to help quickly solve your problem.
aclnnFakeQuantPerTensorAffineCachemask 在构造 scale/zeroPoint 的广播 shape 时, 直接取 self 第0维大小: int64_t tensorSize = (int64_t)(selfContiguous->GetViewShape().GetDim(0)); 当 self 为0维(标量)张量时,GetDimNum() 为 0,GetDim(0) 属于越界访问, 导致接口报错/异常,无法支持标量输入(PyTorch 原生 fake_quantize_per_tensor_affine_cachemask 支持标量)。
先获取维度数 dimNum,当 dimNum == 0 时 tensorSize 取 1(标量张量元素个数为1), 否则取第0维大小: int64_t dimNum = selfContiguous->GetViewShape().GetDimNum(); int64_t tensorSize = (dimNum == 0) ? (int64_t)1 : (int64_t)(selfContiguous->GetViewShape().GetDim(0)); 保证 scale/zeroPoint 广播到合法的 [1] shape,标量场景正常走 FakeQuantAffineCachemask 计算链路。
仅影响 aclnnFakeQuantPerTensorAffineCachemaskGetWorkspaceSize 中 expectShape 的构造逻辑,非0维输入行为不变。
补充0维标量输入的 UT/ST 用例,验证输出和 mask 与期望一致; 回归原有非0维 shape 用例,结果无变化
Ascend950PR
正常计算,标量作为shape为(1,)的tensor。
结果报错
/assign
Thanks for sending an issue! Please fill in the following template to help quickly solve your problem.
Describe the current behavior / 问题描述 (Mandatory / 必填)
问题背景
aclnnFakeQuantPerTensorAffineCachemask 在构造 scale/zeroPoint 的广播 shape 时,
直接取 self 第0维大小:
int64_t tensorSize = (int64_t)(selfContiguous->GetViewShape().GetDim(0));
当 self 为0维(标量)张量时,GetDimNum() 为 0,GetDim(0) 属于越界访问,
导致接口报错/异常,无法支持标量输入(PyTorch 原生
fake_quantize_per_tensor_affine_cachemask 支持标量)。
修改方案
先获取维度数 dimNum,当 dimNum == 0 时 tensorSize 取 1(标量张量元素个数为1),
否则取第0维大小:
int64_t dimNum = selfContiguous->GetViewShape().GetDimNum();
int64_t tensorSize = (dimNum == 0) ? (int64_t)1
: (int64_t)(selfContiguous->GetViewShape().GetDim(0));
保证 scale/zeroPoint 广播到合法的 [1] shape,标量场景正常走
FakeQuantAffineCachemask 计算链路。
影响范围
仅影响 aclnnFakeQuantPerTensorAffineCachemaskGetWorkspaceSize 中
expectShape 的构造逻辑,非0维输入行为不变。
自测
补充0维标量输入的 UT/ST 用例,验证输出和 mask 与期望一致;
回归原有非0维 shape 用例,结果无变化
Environment / 环境信息 (Mandatory / 必填)
Ascend950PR
Steps to reproduce the issue / 重现步骤 (Mandatory / 必填)
aclnnFakeQuantPerTensorAffineCachemask 在构造 scale/zeroPoint 的广播 shape 时,
直接取 self 第0维大小:
int64_t tensorSize = (int64_t)(selfContiguous->GetViewShape().GetDim(0));
当 self 为0维(标量)张量时,GetDimNum() 为 0,GetDim(0) 属于越界访问,
导致接口报错/异常,无法支持标量输入(PyTorch 原生
fake_quantize_per_tensor_affine_cachemask 支持标量)。
Describe the expected behavior / 预期结果 (Mandatory / 必填)
正常计算,标量作为shape为(1,)的tensor。
Related log / screenshot / 日志 / 截图 (Mandatory / 必填)
结果报错
Special notes for this issue/备注 (Optional / 选填)