77
88from typing import TYPE_CHECKING , ClassVar , final
99
10- from pydantic import PositiveInt , model_validator
10+ from pydantic import BaseModel , ConfigDict , Field , PositiveInt , model_validator
1111
1212from coreai_opt .config import (
1313 CompressionConfig ,
3131_PALETTIZATION_SPEC = "palettization_spec"
3232
3333
34+ class PATSchedule (BaseModel ):
35+ """Schedule for enabling palettization-aware training (PAT).
36+
37+ Defines the step threshold at which a module's fake palettization
38+ forward pass becomes active. Used with ``KMeansPalettizer.step()``.
39+
40+ Attributes:
41+ enable_fake_palettize: Step count at which fake palettization is
42+ enabled. Must be >= 0.
43+
44+ Example:
45+ >>> schedule = PATSchedule(enable_fake_palettize=500)
46+ """
47+
48+ model_config = ConfigDict (frozen = True )
49+
50+ enable_fake_palettize : int = Field (default = 0 , ge = 0 )
51+
52+ def _compute_state (self , step_count : int ) -> bool :
53+ """Return whether fake palettization should be active at the given step."""
54+ return step_count >= self .enable_fake_palettize
55+
56+
3457class OpKMeansPalettizerConfig (WeightOnlyOpValidationMixin , OpCompressionConfig [PalettizationSpec ]):
3558 """
3659 Configuration class for palettization at the operation level.
@@ -140,6 +163,11 @@ class ModuleKMeansPalettizerConfig(
140163 K-means clustering. Higher values preserve more precision but may reduce
141164 speed benefits. Only used when enable_fast_kmeans_mode is True. Default: 4.
142165
166+ pat_schedule (PATSchedule | None): Schedule controlling when this
167+ module's palettization is active during a training_mode() loop.
168+ If None, palettization is active immediately once training_mode()
169+ begins. Default: None.
170+
143171 Example:
144172 >>> config = ModuleKMeansPalettizerConfig() # Uses defaults
145173 >>> # Or with custom settings:
@@ -160,6 +188,7 @@ def __init_subclass__(cls, **kwargs):
160188
161189 enable_fast_kmeans_mode : bool = True
162190 rounding_precision : PositiveInt = 4
191+ pat_schedule : PATSchedule | None = None
163192
164193 # Namespace exposing built-in preset constructors.
165194 presets : ClassVar [_ModuleKMeansPalettizerConfigPresets ]
@@ -183,6 +212,12 @@ def validate_fast_kmeans_cluster_dim_constraint(
183212
184213 return self
185214
215+ def _get_fake_module_kwargs (self ) -> dict :
216+ """Exclude palettizer-only fields from the fake-palettize constructor args."""
217+ base = super ()._get_fake_module_kwargs ()
218+ base .pop ("pat_schedule" , None )
219+ return base
220+
186221
187222@final
188223class KMeansPalettizerConfig (CompressionConfig [ModuleKMeansPalettizerConfig ]):
0 commit comments