77
88from __future__ import annotations
99
10- from abc import abstractmethod
11- from typing import TYPE_CHECKING , Any
10+ from abc import ABC , abstractmethod
11+ from typing import TYPE_CHECKING , Any , ClassVar
1212
1313import torch
1414from pydantic import BaseModel , ConfigDict , model_serializer
1515
16- from coreai_opt ._utils .registry_utils import (
17- ClassRegistryMixin as _ClassRegistryMixin ,
18- ConfigRegistryMixin as _ConfigRegistryMixin ,
19- )
16+ from coreai_opt ._utils .registry_utils import ConfigRegistryMixin as _ConfigRegistryMixin
2017
2118if TYPE_CHECKING :
2219 from coreai_opt .palettization .kmeans .kmeans_fake_palettize import _KMeansFakePalettize
2320
2421
25- class TrainingStrategy (_ClassRegistryMixin ):
22+ class TrainingStrategy (ABC ):
2623 """Contract for a fake-palettize module's training-time forward pass."""
2724
2825 @abstractmethod
@@ -38,7 +35,6 @@ def train_forward(self, module: _KMeansFakePalettize, weight: torch.Tensor) -> t
3835 raise NotImplementedError
3936
4037
41- @TrainingStrategy .register ("default" )
4238class _DefaultTrainingStrategy (TrainingStrategy ):
4339 """Post-training, one-shot k-means — today's KMeansPalettizer behavior.
4440
@@ -61,13 +57,16 @@ def train_forward(self, module: _KMeansFakePalettize, weight: torch.Tensor) -> t
6157class TrainingStrategyConfig (BaseModel , _ConfigRegistryMixin ):
6258 """Base class for a fake-palettize module's training-strategy settings.
6359
64- Each subclass is registered under the same key as its paired
65- ``TrainingStrategy`` behavior class (e.g. ``"default"``) ; ``build_strategy()``
66- resolves and constructs that pair from this config's own fields.
60+ Each subclass points ``_strategy_cls`` at its paired ``TrainingStrategy``
61+ behavior class; ``build_strategy()`` constructs that strategy from this
62+ config's own fields.
6763 """
6864
6965 model_config = ConfigDict (frozen = True , extra = "forbid" )
7066
67+ # Each subclass points this at its paired TrainingStrategy behavior class.
68+ _strategy_cls : ClassVar [type [TrainingStrategy ]]
69+
7170 @model_serializer
7271 def _serialize_model (self ) -> dict [str , Any ]:
7372 """Custom serializer that includes the registry type."""
@@ -91,15 +90,12 @@ def _serialize_model(self) -> dict[str, Any]:
9190
9291 def build_strategy (self ) -> TrainingStrategy :
9392 """Construct this config's paired ``TrainingStrategy`` behavior instance."""
94- for key , registered_class in TrainingStrategyConfig .REGISTRY .items ():
95- if registered_class is type (self ):
96- kwargs = {k : v for k , v in self .model_dump ().items () if k != "type" }
97- return TrainingStrategy .resolve (key )(** kwargs )
98- raise ValueError (
99- f"{ type (self ).__name__ } is not registered in TrainingStrategyConfig's registry."
100- )
93+ kwargs = {k : v for k , v in self .model_dump ().items () if k != "type" }
94+ return self ._strategy_cls (** kwargs )
10195
10296
10397@TrainingStrategyConfig .register ("default" )
10498class DefaultTrainingConfig (TrainingStrategyConfig ):
10599 """Settings for the default, post-training one-shot k-means strategy. No fields."""
100+
101+ _strategy_cls = _DefaultTrainingStrategy
0 commit comments