Skip to content

Commit c168890

Browse files
jing.songPostMalone1998
authored andcommitted
1.support scatterElements when update is weightop
2.support InstanceNorm, RMSNorm, clip, pow for mars3 3.update mars3 backend for softmax_log global bugfix 4.support softplus bf16 Change-Id: I77bef269d12a518bd17f88c13264d6caf0aa866b
1 parent 98ac538 commit c168890

10 files changed

Lines changed: 45 additions & 20 deletions

File tree

lib/Conversion/TopToTpu/BM1684X/Clip.cpp

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -40,7 +40,10 @@ void ClipLowering::LoweringINT4(PatternRewriter &rewriter, top::ClipOp op,
4040
void ClipLowering::LoweringINT8(PatternRewriter &rewriter, top::ClipOp op,
4141
bool asymmetric) const {
4242
// nodechip fix8b to be implemented,
43-
LoweringF16(rewriter, op);
43+
if(module::isMARS3())
44+
LoweringBF16(rewriter, op);
45+
else
46+
LoweringF16(rewriter, op);
4447
}
4548

4649
void ClipLowering::LoweringBF16(PatternRewriter &rewriter, top::ClipOp op) const {

lib/Conversion/TopToTpu/BM1684X/InstanceNorm.cpp

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -60,7 +60,10 @@ void InstanceNormLowering::LoweringF32(PatternRewriter &rewriter,
6060
void InstanceNormLowering::LoweringINT8(PatternRewriter &rewriter,
6161
top::InstanceNormOp op,
6262
bool asymmetric) const {
63-
LoweringInstanceNorm(rewriter, op, rewriter.getF32Type());
63+
if(module::isMARS3())
64+
LoweringInstanceNorm(rewriter, op, rewriter.getBF16Type());
65+
else
66+
LoweringInstanceNorm(rewriter, op, rewriter.getF32Type());
6467
}
6568

6669
void InstanceNormLowering::LoweringINT4(PatternRewriter &rewriter,
@@ -71,7 +74,7 @@ void InstanceNormLowering::LoweringINT4(PatternRewriter &rewriter,
7174

7275
void InstanceNormLowering::LoweringBF16(PatternRewriter &rewriter,
7376
top::InstanceNormOp op) const {
74-
LoweringInstanceNorm(rewriter, op, rewriter.getF32Type());
77+
LoweringInstanceNorm(rewriter, op, rewriter.getBF16Type());
7578
}
7679

7780
void InstanceNormLowering::LoweringF16(PatternRewriter &rewriter,

lib/Conversion/TopToTpu/BM1684X/Pow.cpp

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -168,20 +168,20 @@ void PowLowering::LoweringBF16(PatternRewriter &rewriter, top::PowOp op) const {
168168
op.replaceAllUsesWith(ex_op.getOperation());
169169
return ex_op.getOutput();
170170
};
171-
// support f32, change type when needed
172171
auto insert_abs = [&rewriter](top::PowOp &op) -> tpu::ActiveOp {
173172
auto name = module::getName(op.getOutput());
174173
std::vector<NamedAttribute> attrs;
175174
auto abs_loc = NameLoc::get(rewriter.getStringAttr(name.str() + "_abs"));
176175
attrs.push_back(rewriter.getNamedAttr(
177176
"mode",
178177
tpu::ActiveModeAttr::get(op.getContext(), tpu::ActiveMode::ABSVAL)));
178+
auto new_type = getQuantBF16Type(op.getResult());
179179
auto abs_op = rewriter.create<tpu::ActiveOp>(
180-
abs_loc, op.getOutput().getType(), ValueRange{op.getInput()}, attrs);
180+
abs_loc, new_type, ValueRange{op.getInput()}, attrs);
181181
op->setOperand(0, abs_op.getOutput());
182182
return abs_op;
183183
};
184-
184+
// support f32, change type when needed
185185
double exponent = op.getExponent().convertToDouble();
186186

187187
if (fmod(exponent, 2) == 0) {

lib/Conversion/TopToTpu/BM1684X/RMSNorm.cpp

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,10 @@ void RMSNormLowering::LoweringF32(PatternRewriter &rewriter,
5858

5959
void RMSNormLowering::LoweringINT8(PatternRewriter &rewriter, top::RMSNormOp op,
6060
bool asymmetric) const {
61-
LoweringF16(rewriter, op);
61+
if(module::isMARS3())
62+
LoweringBF16(rewriter, op);
63+
else
64+
LoweringF16(rewriter, op);
6265
}
6366

6467
void RMSNormLowering::LoweringINT4(PatternRewriter &rewriter, top::RMSNormOp op,

lib/Conversion/TopToTpu/BM1684X/ScatterElements.cpp

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -98,8 +98,11 @@ void ScatterElementsLowering::LoweringINT8(PatternRewriter &rewriter,
9898
// lowering_common_int8<tpu::ScatterElementsOp>(rewriter, op.getOperation(),
9999
// asymmetric);
100100
// Please implent lowering quant for weight if necessary
101-
if(module::isWeight(op.getInput())){
102-
LoweringF32(rewriter, op);
101+
if(module::isWeight(op.getInput()) || module::isWeight(op.getIndices()) || module::isWeight(op.getUpdates())){
102+
if(module::isMARS3())
103+
LoweringBF16(rewriter, op);
104+
else
105+
LoweringF32(rewriter, op);
103106
return;
104107
}
105108
auto new_type = getQuantInt8Type(op.getOutput());

lib/Conversion/TopToTpu/BM1684X/Softplus.cpp

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,12 +37,25 @@ void SoftplusLowering::LoweringINT8(PatternRewriter &rewriter,
3737

3838
void SoftplusLowering::LoweringBF16(PatternRewriter &rewriter,
3939
top::SoftplusOp op) const {
40-
LoweringF32(rewriter, op);
40+
if(module::isMARS3()){
41+
auto op_ = op.getOperation();
42+
op_->setAttr("mode", tpu::ActiveModeAttr::get(op.getContext(),
43+
tpu::ActiveMode::SOFT_PLUS));
44+
lowering_common_bf16<tpu::ActiveOp>(rewriter, op_);
45+
} else {
46+
LoweringF32(rewriter, op);
47+
}
48+
4149
}
4250

4351
void SoftplusLowering::LoweringF16(PatternRewriter &rewriter,
4452
top::SoftplusOp op) const {
4553
LoweringF32(rewriter, op);
54+
// uncomment when needed
55+
// auto op_ = op.getOperation();
56+
// op_->setAttr("mode", tpu::ActiveModeAttr::get(op.getContext(),
57+
// tpu::ActiveMode::SOFT_PLUS));
58+
// lowering_common_f16<tpu::ActiveOp>(rewriter, op_);
4659
}
4760

4861
void SoftplusLowering::LoweringF8(PatternRewriter &rewriter,

python/test/test_onnx.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -186,7 +186,7 @@ def __init__(self,
186186
"ShapeSlice": (self.test_ShapeSlice, Y, N, N, N, N, Y),
187187
"SiLU": (self.test_SiLU, Y, Y, Y, Y, Y, Y),
188188
"Softmax": (self.test_Softmax, Y, Y, Y, Y, Y, Y),
189-
"Softplus": (self.test_Softplus, Y, Y, Y, Y, Y, N),
189+
"Softplus": (self.test_Softplus, Y, Y, Y, Y, Y, Y),
190190
"Space2Depth": (self.test_Space2Depth, Y, Y, Y, N, Y, Y),
191191
"Squeeze": (self.test_Squeeze, Y, Y, Y, Y, Y, Y),
192192
"Sigmoid": (self.test_Sigmoid, Y, Y, Y, Y, Y, Y),
@@ -240,23 +240,23 @@ def __init__(self,
240240
"TorchGroupNorm2": (self.test_TorchGroupNorm2, Y, Y, Y, N, Y, Y),
241241
"TorchGRU": (self.test_TorchGRU, N, Y, Y, Y, Y, N),
242242
"TorchIdentity": (self.test_TorchIdentity, Y, Y, Y, Y, Y, Y),
243-
"TorchIndexCopy": (self.test_TorchIndexCopy, N, N, N, N, N, N),
244-
"TorchInstanceNorm": (self.test_TorchInstanceNorm, N, Y, Y, N, Y, N),
245-
"TorchInstanceNorm2": (self.test_TorchInstanceNorm2, N, Y, Y, N, Y, N),
243+
"TorchIndexCopy": (self.test_TorchIndexCopy, N, N, N, N, N, Y),
244+
"TorchInstanceNorm": (self.test_TorchInstanceNorm, N, Y, Y, N, Y, Y),
245+
"TorchInstanceNorm2": (self.test_TorchInstanceNorm2, N, Y, Y, N, Y, Y),
246246
"TorchLayerGroup": (self.test_TorchLayerGroup, Y, Y, Y, Y, Y, Y),
247247
"TorchLayerNorm": (self.test_TorchLayerNorm, Y, Y, Y, Y, Y, Y),
248-
"TorchLayerNorm2": (self.test_TorchLayerNorm2, Y, Y, Y, Y, Y, N),
249-
"TorchLogSoftmax": (self.test_TorchLogSoftmax, Y, Y, Y, Y, Y, N),
248+
"TorchLayerNorm2": (self.test_TorchLayerNorm2, Y, Y, Y, Y, Y, Y),
249+
"TorchLogSoftmax": (self.test_TorchLogSoftmax, Y, Y, Y, Y, Y, Y),
250250
"TorchLSTM": (self.test_TorchLSTM, Y, Y, Y, Y, Y, N),
251251
"TorchMaskedFill": (self.test_TorchMaskedFill, N, Y, Y, N, Y, Y),
252252
"TorchNonZero": (self.test_TorchNonZero, N, Y, Y, N, Y, N),
253-
"TorchNormalize": (self.test_TorchNormalize, N, Y, Y, N, N, N),
253+
"TorchNormalize": (self.test_TorchNormalize, N, Y, Y, N, N, Y),
254254
"TorchReflectionPad": (self.test_TorchReflectionPad, N, Y, Y, Y, Y, Y),
255-
"TorchRMSNorm": (self.test_TorchRMSNorm, N, Y, Y, N, Y, N),
255+
"TorchRMSNorm": (self.test_TorchRMSNorm, N, Y, Y, N, Y, Y),
256256
"TorchRoiAlign": (self.test_TorchRoiAlign, N, Y, Y, N, Y, N),
257257
"TorchScatterND": (self.test_TorchScatterND, N, Y, Y, Y, Y, N),
258258
"TorchSize": (self.test_TorchSize, Y, Y, Y, Y, Y, Y),
259-
"TorchStd": (self.test_TorchStd, N, Y, Y, Y, Y, N),
259+
"TorchStd": (self.test_TorchStd, N, Y, Y, Y, Y, Y),
260260
"TorchWhere": (self.test_TorchWhere, N, Y, Y, N, Y, Y),
261261
"TorchZeroPad": (self.test_TorchZeroPad, N, Y, Y, Y, Y, Y),
262262
#########################################

third_party/nntoolchain/README.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -57,7 +57,7 @@ elease: unset EXTRA_CONFIG && rebuild_backend_lib_release_cmodel
5757
cp build/backend_api/libbackend_sg2380.so /workspace/tpu-mlir/third_party/nntoolchain/lib/libbackend_sg2380.so
5858
cp build_runtime/firmware_core/libcmodel_firmware.so /workspace/tpu-mlir/third_party/nntoolchain/lib/libcmodel_sg2380.so
5959

60-
#mars3 sha256: e5b864414e3a4feb8a343f7c3417601addb6cc82
60+
#mars3 sha256: fde8a835acd0a4c194b02d3af836c7c300f849b9
6161
cd TPU1686
6262
source scripts/envsetup.sh mars3
6363
debug: rebuild_backend_lib_cmodel
1.16 MB
Binary file not shown.
1.01 MB
Binary file not shown.

0 commit comments

Comments
 (0)