Skip to content

Commit 92a206f

Browse files
committed
Remove registration for TrainingStrategy
1 parent a713e2b commit 92a206f

3 files changed

Lines changed: 27 additions & 28 deletions

File tree

src/coreai_opt/palettization/spec/spec.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -64,9 +64,11 @@ class PalettizationSpec(CompressionSpec):
6464
training_strategy_config: Which training-time behavior this weight's
6565
fake-palettize module uses, and that strategy's settings.
6666
``DefaultTrainingConfig()`` is post-training, one-shot k-means
67-
(today's KMeansPalettizer behavior). Additional strategies can be
68-
registered via ``TrainingStrategy.register()`` /
69-
``TrainingStrategyConfig.register()``. Default: DefaultTrainingConfig().
67+
(today's KMeansPalettizer behavior). Additional strategies are added
68+
by subclassing ``TrainingStrategyConfig`` (registered via
69+
``TrainingStrategyConfig.register()``) and pointing its
70+
``_strategy_cls`` at a ``TrainingStrategy`` subclass. Default:
71+
DefaultTrainingConfig().
7072
7173
Example:
7274
>>> # Basic 4-bit palettization

src/coreai_opt/palettization/spec/training_strategy.py

Lines changed: 14 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -7,22 +7,19 @@
77

88
from __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

1313
import torch
1414
from 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

2118
if 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")
4238
class _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
6157
class 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")
10498
class DefaultTrainingConfig(TrainingStrategyConfig):
10599
"""Settings for the default, post-training one-shot k-means strategy. No fields."""
100+
101+
_strategy_cls = _DefaultTrainingStrategy

tests/palettization/test_kmeans_fake_palettize.py

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1611,18 +1611,19 @@ def test_lut_consistent_with_quantized_lut_and_scale(self):
16111611

16121612

16131613
class TestTrainingStrategy:
1614-
"""The training-strategy registry and its paired config objects.
1614+
"""The training-strategy config registry and its paired behavior classes.
16151615
16161616
Scope is the generic OSS machinery (default strategy only); concrete
1617-
strategies registered elsewhere are tested with those strategies.
1617+
strategies defined elsewhere are tested with those strategies.
16181618
"""
16191619

1620-
def test_resolve_default(self):
1621-
assert TrainingStrategy.resolve("default") is _DefaultTrainingStrategy
1620+
def test_default_config_points_at_default_strategy(self):
1621+
assert DefaultTrainingConfig._strategy_cls is _DefaultTrainingStrategy
1622+
assert issubclass(DefaultTrainingConfig._strategy_cls, TrainingStrategy)
16221623

1623-
def test_resolve_unregistered_raises(self):
1624-
with pytest.raises(ValueError):
1625-
TrainingStrategy.resolve("nonexistent")
1624+
def test_build_from_dict_unregistered_raises(self):
1625+
with pytest.raises(KeyError):
1626+
TrainingStrategyConfig.maybe_build_from_dict({"type": "nonexistent"})
16261627

16271628
def test_default_config_builds_default_strategy(self):
16281629
assert isinstance(DefaultTrainingConfig().build_strategy(), _DefaultTrainingStrategy)

0 commit comments

Comments
 (0)