@@ -604,9 +604,9 @@ def test_partial_vs_direct_instantiation(self):
604604 assert fq_partial_out .dtype == x .dtype
605605 assert fq_direct_out .dtype == x .dtype
606606
607- def test_update_partial_qparams_calculator_single_attr (self ):
608- """
609- update_partial_qparams_calculator overrides one attribute on each constructed calculator .
607+ def test_reconstruct_partial_qparams_calculator (self ):
608+ """reconstruct_partial_qparams_calculator overrides attributes without mutating the original
609+ partial .
610610 """
611611 spec = QuantizationSpec (
612612 dtype = torch .int8 ,
@@ -619,78 +619,33 @@ def test_update_partial_qparams_calculator_single_attr(self):
619619 spec , quantization_target = CompressionTargetTensor .WEIGHT
620620 )
621621
622- updated = QuantizationComponentFactory .update_partial_qparams_calculator (
622+ # Overriding a single attribute updates the constructed calculator.
623+ updated_single = QuantizationComponentFactory .reconstruct_partial_qparams_calculator (
623624 partial , qscheme = QuantizationScheme .ASYMMETRIC
624625 )
625-
626- fq = updated ()
626+ fq = updated_single ()
627627 assert fq .qparams_calculator .qscheme == QuantizationScheme .ASYMMETRIC
628628 # qscheme property on the fake quantizer delegates to the calculator
629629 assert fq .qscheme == QuantizationScheme .ASYMMETRIC
630630
631- def test_update_partial_qparams_calculator_multiple_attrs (self ):
632- """update_partial_qparams_calculator overrides multiple attributes in one call."""
633- spec = QuantizationSpec (
634- dtype = torch .int8 ,
635- qscheme = "symmetric" ,
636- granularity = PerTensorGranularity (),
637- qparam_calculator_cls = StaticQParamsCalculator ,
638- range_calculator_cls = MinMaxRangeCalculator ,
639- )
640- partial = QuantizationComponentFactory .create_fake_quantizer_partial (
641- spec , quantization_target = CompressionTargetTensor .WEIGHT
642- )
643-
644- updated = QuantizationComponentFactory .update_partial_qparams_calculator (
631+ # Overriding multiple attributes in one call updates all of them.
632+ updated_multiple = QuantizationComponentFactory .reconstruct_partial_qparams_calculator (
645633 partial ,
646634 qscheme = QuantizationScheme .ASYMMETRIC ,
647635 float_range = (0.0 , 1.0 ),
648636 )
649-
650- fq = updated ()
637+ fq = updated_multiple ()
651638 assert fq .qparams_calculator .qscheme == QuantizationScheme .ASYMMETRIC
652639 assert fq .qparams_calculator .float_range == (0.0 , 1.0 )
653640
654- def test_update_partial_qparams_calculator_does_not_mutate_original (self ):
655- """update_partial_qparams_calculator returns a new partial; the original is unchanged."""
656- spec = QuantizationSpec (
657- dtype = torch .int8 ,
658- qscheme = "symmetric" ,
659- granularity = PerTensorGranularity (),
660- qparam_calculator_cls = StaticQParamsCalculator ,
661- range_calculator_cls = MinMaxRangeCalculator ,
662- )
663- partial = QuantizationComponentFactory .create_fake_quantizer_partial (
664- spec , quantization_target = CompressionTargetTensor .WEIGHT
665- )
666-
667- _ = QuantizationComponentFactory .update_partial_qparams_calculator (
668- partial , qscheme = QuantizationScheme .ASYMMETRIC
669- )
670-
671- # Original partial still produces calculators with the original qscheme.
641+ # The original partial is unaffected by either override.
672642 fq_original = partial ()
673643 assert fq_original .qparams_calculator .qscheme == QuantizationScheme .SYMMETRIC
674644
675- def test_update_partial_qparams_calculator_independent_instances (self ):
676- """Each call to an updated partial creates an independent calculator with the override."""
677- spec = QuantizationSpec (
678- dtype = torch .int8 ,
679- qscheme = "symmetric" ,
680- granularity = PerTensorGranularity (),
681- qparam_calculator_cls = MovingAverageQParamsCalculator ,
682- range_calculator_cls = MinMaxRangeCalculator ,
683- )
684- partial = QuantizationComponentFactory .create_fake_quantizer_partial (
685- spec , quantization_target = CompressionTargetTensor .ACTIVATION
686- )
687- updated = QuantizationComponentFactory .update_partial_qparams_calculator (
688- partial , float_range = (0.0 , 1.0 )
689- )
690-
691- fq1 = updated ()
692- fq2 = updated ()
693-
645+ # Each call to an updated partial creates an independent calculator with the override
646+ # applied.
647+ fq1 = updated_multiple ()
648+ fq2 = updated_multiple ()
694649 assert id (fq1 .qparams_calculator ) != id (fq2 .qparams_calculator )
695650 assert fq1 .qparams_calculator .float_range == (0.0 , 1.0 )
696651 assert fq2 .qparams_calculator .float_range == (0.0 , 1.0 )
0 commit comments