diff --git a/tensilelite/Tensile/Components/GlobalWriteBatch.py b/tensilelite/Tensile/Components/GlobalWriteBatch.py index cf97ca9adc..74fc7d6c45 100644 --- a/tensilelite/Tensile/Components/GlobalWriteBatch.py +++ b/tensilelite/Tensile/Components/GlobalWriteBatch.py @@ -2012,7 +2012,7 @@ def _addSumAlphaWithCBeta(self, kernel, ss, gwvw, elementIdx, vc0, tmpVgpr, cvtV # Generate single f32 code if edge is detected. isPK = False if ((vi + 1) == self.gwvw) and ((self.gwvw % 2) == 1): - if self.parentWriter.states.archCaps["NoSDWA"]: #cm review + if self.parentWriter.states.archCaps["NoSDWA"]: sb = 0 if self.gwvw == 1 else 1 module.add(VCvtFP8toF32(dst=vgpr(tmpVgpr), src=vgpr(dataV), vop3=VOP3PModifiers(op_sel=[0,sb]))) else: @@ -2023,7 +2023,7 @@ def _addSumAlphaWithCBeta(self, kernel, ss, gwvw, elementIdx, vc0, tmpVgpr, cvtV continue else: isPK = True - if self.parentWriter.states.archCaps["NoSDWA"]: #cm review + if self.parentWriter.states.archCaps["NoSDWA"]: sb = 0 if vi ==0 else 1 module.add(VCvtPkFP8toF32(dst=vgpr(tmpVgpr, 2), src=vgpr(dataV), vop3=VOP3PModifiers(op_sel=[sb]))) else: @@ -2041,9 +2041,9 @@ def _addSumAlphaWithCBeta(self, kernel, ss, gwvw, elementIdx, vc0, tmpVgpr, cvtV # Generate single f32 code if edge is detected. isPK = False if ((vi + 1) == self.gwvw) and ((self.gwvw % 2) == 1): - if self.parentWriter.states.archCaps["NoSDWA"]: #cm review + if self.parentWriter.states.archCaps["NoSDWA"]: sb = 0 if self.gwvw == 1 else 1 - module.add(VCvtFP8toF32(dst=vgpr(tmpVgpr), src=vgpr(dataV), vop3=VOP3PModifiers(op_sel=[0,sb]))) + module.add(VCvtBF8toF32(dst=vgpr(tmpVgpr), src=vgpr(dataV), vop3=VOP3PModifiers(op_sel=[0,sb]))) else: sb = SelectBit.BYTE_0 if self.gwvw == 1 else SelectBit.BYTE_2 module.add(VCvtBF8toF32(dst=vgpr(tmpVgpr), src=vgpr(dataV), sdwa=SDWAModifiers(src0_sel=sb))) @@ -2052,9 +2052,9 @@ def _addSumAlphaWithCBeta(self, kernel, ss, gwvw, elementIdx, vc0, tmpVgpr, cvtV continue else: isPK = True - if self.parentWriter.states.archCaps["NoSDWA"]: #cm review + if self.parentWriter.states.archCaps["NoSDWA"]: sb = 0 if vi ==0 else 1 - module.add(VCvtPkFP8toF32(dst=vgpr(tmpVgpr, 2), src=vgpr(dataV), vop3=VOP3PModifiers(op_sel=[sb]))) + module.add(VCvtPkBF8toF32(dst=vgpr(tmpVgpr, 2), src=vgpr(dataV), vop3=VOP3PModifiers(op_sel=[sb]))) else: sb = SelectBit.WORD_0 if vi == 0 else SelectBit.WORD_1 module.add(VCvtPkBF8toF32(dst=vgpr(tmpVgpr, 2), src=vgpr(dataV), sdwa=SDWAModifiers(src0_sel=sb))) diff --git a/tensilelite/Tensile/TensileInstructions/Instructions.py b/tensilelite/Tensile/TensileInstructions/Instructions.py index ee60729cad..c535d2b379 100644 --- a/tensilelite/Tensile/TensileInstructions/Instructions.py +++ b/tensilelite/Tensile/TensileInstructions/Instructions.py @@ -2527,8 +2527,8 @@ def __init__(self, dst, src, sdwa: Optional[SDWAModifiers] = None, vop3: Optiona self.setInst("v_cvt_f32_fp8") class VCvtBF8toF32(VCvtInstruction): - def __init__(self, dst, src, sdwa: Optional[SDWAModifiers] = None, comment="") -> None: - super().__init__(CvtType.CVT_BF8_to_F32, dst, src, sdwa, None, comment) + def __init__(self, dst, src, sdwa: Optional[SDWAModifiers] = None, vop3: Optional[VOP3PModifiers] = None, comment="") -> None: + super().__init__(CvtType.CVT_BF8_to_F32, dst, src, sdwa, vop3, comment) self.setInst("v_cvt_f32_bf8") class VCvtPkFP8toF32(VCvtInstruction): @@ -2537,8 +2537,8 @@ def __init__(self, dst, src, sdwa: Optional[SDWAModifiers] = None, vop3: Optiona self.setInst("v_cvt_pk_f32_fp8") class VCvtPkBF8toF32(VCvtInstruction): - def __init__(self, dst, src, sdwa: Optional[SDWAModifiers] = None, comment="") -> None: - super().__init__(CvtType.CVT_PK_BF8_to_F32, dst, src, sdwa, None, comment) + def __init__(self, dst, src, sdwa: Optional[SDWAModifiers] = None, vop3: Optional[VOP3PModifiers] = None, comment="") -> None: + super().__init__(CvtType.CVT_PK_BF8_to_F32, dst, src, sdwa, vop3, comment) self.setInst("v_cvt_pk_f32_bf8") class VCvtPkF32toFP8(VCvtInstruction):