Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions docs/src/palettization/overview.md
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,24 @@ Note that `Sensitive K-means` shares the same `KMeansPalettizer` and `KMeansPale

For more details on how to use {class}`~coreai_opt.palettization.config.KMeansPalettizerConfig`, {class}`~coreai_opt.palettization.config.ModuleKMeansPalettizerConfig` to apply different settings to different weights in the model, see [Palettization Config](config.md).

## Loading a Prepared Model

Preparing a model runs k-means clustering to compute each layer's codebook, which can be expensive for large models. If you have already prepared a model and saved its `state_dict`, reload it into a fresh palettizer with `prepare(..., state_dict=...)` to skip clustering:

```python
# ---- first run: prepare and save ----
palettizer = KMeansPalettizer(model, config)
prepared_model = palettizer.prepare(example_inputs)
torch.save(prepared_model.state_dict(), "palettized_state.pt")

# ---- later run: reload without re-clustering ----
state_dict = torch.load("palettized_state.pt", weights_only=True)
palettizer = KMeansPalettizer(MyModel().eval(), config) # same config as the saved run
prepared_model = palettizer.prepare(example_inputs, state_dict=state_dict)
```

The `config` and model architecture must match the run that produced the state dict; a mismatch raises an error rather than loading a wrong codebook. Sensitivities are not part of the state dict — reload them separately via `prepare(..., sensitivity_path=...)` if a later re-clustering step needs them.

## Training a Palettized model

A `KMeansPalettizer` palettized model can still be fine-tuned in a training pipeline. As palettization is a hard assignment lookup, gradients cannot be propagated for palettized weights, meaning any parameter which is palettized will not update during `optimizer.step()` (the palettization codebook and index assignments will also be fixed).
Expand Down
52 changes: 49 additions & 3 deletions src/coreai_opt/palettization/kmeans/kmeans_fake_palettize.py
Original file line number Diff line number Diff line change
Expand Up @@ -197,6 +197,24 @@ def ensure_initialized(self, tensor: torch.Tensor) -> None:
)
self._disabled = True

def check_compatible(self, weight: torch.Tensor) -> bool:
"""Return whether ``weight``'s shape is compatible with this module's spec.

Runs the reshape/block prefix on a meta (zero-storage) copy of ``weight``:
no data is read, no scaling is applied, and nothing is allocated.

Args:
weight (torch.Tensor): Weight tensor in its original shape.

Returns:
bool: True if the shape is compatible with the configured spec.
"""
try:
self._reshape_and_block(weight.to("meta"))
except (_IncompatibleClusterDimError, _IncompatibleGranularityError):
return False
return True

def forward_enabled(self, tensor: torch.Tensor) -> torch.Tensor:
if self.training:
return self._training_strategy.train_forward(self, tensor)
Expand Down Expand Up @@ -233,17 +251,25 @@ def _maybe_refresh_indices(self, weight: torch.Tensor) -> None:
self.indices = self._assign_indices(weight, self.centroids).detach()
self._indices_stale = False

def _save_to_state_dict(self, destination, prefix, keep_vars):
"""Omit stale indices so they are recomputed from centroids on load."""
super()._save_to_state_dict(destination, prefix, keep_vars)
if self._indices_stale:
destination.pop(prefix + "indices", None)
Comment on lines +257 to +258

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Wouldn't they be overwritten when they are loaded?

I'm wondering - running this as part of coretorch training would run checkpointing - and when we reload the checkpoint between epochs, this would re-compute the indices; would that be unnecessary?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

_indices_stale only gets set to True when it goes out of sync with centroids, in which case it is necessary to recompute them during the next forward pass. I feel like we should only write them out to the state dict when they are valid, or else the state dict itself is in a bad state.

If indices are not stale, they get written out along with centroids so the next load picks up both of them and no further recomputation is needed. But if they are stale, one could either run a forward pass before saving the state dict to recompute them and save both centroids and indices to state dict, or we just write out centroids only and let the first forward pass after loading recompute the indices. Either way that recomputation has to happen at some point.


def _load_from_state_dict(self, state_dict, prefix, *args, **kwargs):
"""Load centroids from a checkpoint, reconstructing them from a legacy
``lut`` buffer when present.
"""
lut_key, centroids_key = prefix + "lut", prefix + "centroids"
lut_key, centroids_key, index_key = prefix + "lut", prefix + "centroids", prefix + "indices"
if centroids_key not in state_dict and lut_key in state_dict:
state_dict[centroids_key] = self._centroids_from_lut(state_dict[lut_key])
had_indices = index_key in state_dict
super()._load_from_state_dict(state_dict, prefix, *args, **kwargs)
if self.centroids is not None:
self._centroids_initialized = True
self._indices_stale = True
scale_ok = not self.enable_per_channel_scale or self.per_channel_scale is not None
self._indices_stale = not (had_indices and scale_ok)

def _centroids_from_lut(self, lut: torch.Tensor) -> torch.Tensor:
"""Invert ``_reshape_lut_tensor`` to recover ``(num_blocks, num_clusters,
Expand Down Expand Up @@ -389,6 +415,24 @@ def _scale_reshape_and_block(self, weight: torch.Tensor) -> tuple[list[torch.Ten
"""
if self.enable_per_channel_scale:
weight = self._scale_by_per_channel_scale(weight)
return self._reshape_and_block(weight)

def _reshape_and_block(self, weight: torch.Tensor) -> tuple[list[torch.Tensor], int]:
"""Reshape to 2D and split into per-block tensors (no scaling).

Shared by the clustering path and the compatibility probe.

Args:
weight (torch.Tensor): Weight tensor in its original shape.

Returns:
tuple[list[torch.Tensor], int]: The per-block 2D tensors and the
resolved palettization axis.

Raises:
_IncompatibleGranularityError: If the shape is incompatible with the granularity.
_IncompatibleClusterDimError: If the shape is incompatible with cluster_dim.
"""
axis = self._resolved_axis

# Produce a 2d tensor with output channel axis remaining as is, and all other axes flattened
Expand Down Expand Up @@ -945,6 +989,8 @@ def _unscale_by_per_channel_scale(self, scaled_weight: torch.Tensor) -> torch.Te
per channel scales.
"""
flattened_scaled_weight = scaled_weight.flatten(1)
flattened_unscaled_weight = flattened_scaled_weight * self.per_channel_scale
flattened_unscaled_weight = flattened_scaled_weight * self.per_channel_scale.to(
flattened_scaled_weight.device
)
unscaled_weight = flattened_unscaled_weight.reshape(scaled_weight.shape)
return unscaled_weight
65 changes: 56 additions & 9 deletions src/coreai_opt/palettization/kmeans/palettizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,6 +154,7 @@ def prepare(
example_inputs: tuple[Any, ...],
sensitivity_path: str | None = None,
num_workers: int = 1,
state_dict: dict[str, torch.Tensor] | None = None,
) -> torch.nn.Module:
"""
Prepare the model for palettization.
Expand All @@ -172,6 +173,9 @@ def prepare(
across layers. It is recommended to use more than one worker
process to parallelize the clustering, especially when multiple
CPUs are available. Defaults to ``1``.
state_dict: Optional state dict from a model previously prepared
with the same config. When provided, buffers are loaded from it
instead of running k-means clustering.

Returns:
The prepared nn.Module with fake palettization
Expand Down Expand Up @@ -208,18 +212,21 @@ def prepare(
sensitivities = torch.load(sensitivity_path, weights_only=True)
self._set_sensitivities_in_fake_palettize_modules(sensitivities)

self._model.apply(_disable_fake_palett)
if state_dict is None:
self._model.apply(_disable_fake_palett)

if self._num_workers > 1:
self._calculate_centroids_parallel(num_workers)
else:
self._calculate_centroids_sequential()
if self._num_workers > 1:
self._calculate_centroids_parallel(num_workers)
else:
self._calculate_centroids_sequential()

# Remove FakePalettize modules that were disabled during the forward
# pass due to incompatible granularity or cluster dimensions.
self._remove_disabled_fake_palett_modules(self._model)
# Remove FakePalettize modules that were disabled during the forward
# pass due to incompatible granularity or cluster dimensions.
self._remove_disabled_fake_palett_modules(self._model)

self._model.apply(_enable_fake_palett)
self._model.apply(_enable_fake_palett)
else:
self._load_prepared_state_dict(state_dict)

# Mark the model as prepared to prevent re-preparation
self._mark_model_as_prepared(prepared_model)
Expand All @@ -229,6 +236,46 @@ def prepare(

return self._model

def _load_prepared_state_dict(self, state_dict: dict[str, torch.Tensor]) -> None:
"""Load a previously-prepared model's buffers instead of computing centroids.

Prunes the same incompatible modules a normal ``prepare()`` would (decided
from config and weight shapes, not from ``state_dict``), strict-loads, then
verifies every surviving module received its centroids.

Args:
state_dict (dict[str, torch.Tensor]): State dict from a model
previously prepared with the same config.

Raises:
RuntimeError: If a compatible module has no centroids after loading,
meaning the checkpoint does not match the palettizer config.
"""
# Disable exactly the modules a normal prepare() would remove, using a
# shape-only (meta) compatibility probe. Structure comes from the config,
# never from the state dict.
for info in self._collect_fake_palett_info(to_cpu=False):
compatible = info.fp_module.check_compatible(info.weight)
info.fp_module._disabled = not compatible
self._remove_disabled_fake_palett_modules(self._model)

self._model.load_state_dict(state_dict)

# Verify all surviving modules have centroids
missing_centroids = []
for info in self._collect_fake_palett_info(to_cpu=False):
fp = info.fp_module
if fp.centroids is None:
missing_centroids.append(fp)

if missing_centroids:
error_msg = (
"Error when loading state dict. The following palettized weights are missing "
"centroids/lut. State dict does not match the palettizer config: "
f"{[fp.tensor_fqn for fp in missing_centroids]}"
)
raise RuntimeError(error_msg)

@contextmanager
def calibration_mode(
self,
Expand Down
60 changes: 54 additions & 6 deletions tests/palettization/test_kmeans_fake_palettize.py
Original file line number Diff line number Diff line change
Expand Up @@ -1860,21 +1860,57 @@ def test_hard_assign_skips_refresh_when_fresh(self):
palettizer.hard_assign(weight)
assert palettizer.indices is indices_before # untouched

def test_load_from_state_dict_marks_initialized_but_stale(self):
"""Loading centroids from a checkpoint marks the module initialized AND
stale (unlike _initialize(), which leaves indices fresh) -- indices are
reconstructed lazily against the loaded centroids on next use.
"""
def test_load_from_state_dict_with_and_without_indices(self):
"""Test _centroids_initialized and _indices_stale settings when loading a state dict with
and without indices."""
spec = PalettizationSpec(n_bits=2, granularity=PerTensorGranularity())
source = _KMeansFakePalettize(**spec.__dict__)
source._initialize(torch.randn(8, 8))
state_dict = source.state_dict()

# Loading state dict with indices sets _indices_stale to False
target = _KMeansFakePalettize(**spec.__dict__)
assert target._centroids_initialized is False
target.load_state_dict(source.state_dict())
target.load_state_dict(state_dict)
assert target._centroids_initialized is True
assert target._indices_stale is False

# Loading state dict without indices sets _indices_stale to True
del state_dict["indices"]
target = _KMeansFakePalettize(**spec.__dict__)
target.load_state_dict(state_dict)
assert target._centroids_initialized is True
assert target._indices_stale is True

def test_load_from_state_dict_missing_scale_is_stale(self):
"""With per-channel scale enabled, loading indices but no scale marks stale
so both are recomputed together on next use.
"""
spec = PalettizationSpec(
n_bits=2, granularity=PerTensorGranularity(), enable_per_channel_scale=True
)
source = _KMeansFakePalettize(**spec.__dict__)
source._initialize(torch.randn(8, 8))
state_dict = source.state_dict()
del state_dict["per_channel_scale"]

target = _KMeansFakePalettize(**spec.__dict__)
target.load_state_dict(state_dict)
assert target._centroids_initialized is True
assert target._indices_stale is True

def test_stale_indices_are_not_saved(self):
"""Stale indices are omitted from the state dict and recomputed on load."""
spec = PalettizationSpec(n_bits=2, granularity=PerTensorGranularity())
source = _KMeansFakePalettize(**spec.__dict__)
source._initialize(torch.randn(8, 8))
source._indices_stale = True
assert "indices" not in source.state_dict()

target = _KMeansFakePalettize(**spec.__dict__)
target.load_state_dict(source.state_dict())
assert target._indices_stale is True


class TestReinitializeOnEnable:
"""enable_fake_palett(reinitialize=True) invalidates cached params on a
Expand Down Expand Up @@ -2219,3 +2255,15 @@ def test_from_cluster_vectors_rejects_element_count_mismatch():
vectors = torch.randn(1, 5, 2) # 5 * 2 = 10 elements, but rows * cols = 8
with pytest.raises(ValueError):
palettizer._from_cluster_vectors(vectors, rows=4, cols=2)


@pytest.mark.parametrize("axis", [0, 1])
@pytest.mark.parametrize("group_size", [3, 4])
def test_check_compatible(axis, group_size):
"""Test check_compatible for a variety of axes and group sizes."""
tensor = torch.randn(8, 8)
spec = PalettizationSpec(
n_bits=2, granularity=PerGroupedChannelGranularity(axis=axis, group_size=group_size)
)
palettizer = _KMeansFakePalettize(**spec.__dict__)
assert palettizer.check_compatible(torch.randn(8, 8)) is (tensor.shape[axis] % group_size == 0)
49 changes: 49 additions & 0 deletions tests/palettization/test_kmeans_palettizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -1646,3 +1646,52 @@ def test_palettize_multihead_attention(

output = prepared_model(simple_mha_model_input)
assert output.shape == (1, 10, 64)


class TestKMeansPalettizerLoadFromStateDict:
"""prepare(state_dict=...) loads a prepared model's buffers instead of clustering."""

def test_reproduces_prepared_model(
self, simple_conv_linear_model, simple_model_input, basic_config, temp_dir
):
"""A skip-loaded model reproduces the source model's forward output exactly."""
src = KMeansPalettizer(copy.deepcopy(simple_conv_linear_model), basic_config)
prepared = src.prepare((simple_model_input,))
state_dict_path = os.path.join(temp_dir, "palettizer_state_dict.pt")
torch.save(prepared.state_dict(), state_dict_path)
with torch.no_grad():
expected = prepared(simple_model_input)

dst = KMeansPalettizer(copy.deepcopy(simple_conv_linear_model), basic_config)
state_dict = torch.load(state_dict_path)
loaded = dst.prepare((simple_model_input,), state_dict=state_dict)
with torch.no_grad():
got = loaded(simple_model_input)

assert torch.equal(expected, got)

def test_skips_clustering(
self, monkeypatch, simple_conv_linear_model, simple_model_input, basic_config
):
"""Loading from a state dict must not run k-means clustering."""
src = KMeansPalettizer(copy.deepcopy(simple_conv_linear_model), basic_config)
state_dict = src.prepare((simple_model_input,)).state_dict()

def _raise_assert(self, *args, **kwargs):
raise AssertionError("clustering must not run when loading from a state dict")

monkeypatch.setattr(_KMeansFakePalettize, "_cluster_to_centroids", _raise_assert)
dst = KMeansPalettizer(copy.deepcopy(simple_conv_linear_model), basic_config)
dst.prepare((simple_model_input,), state_dict=state_dict)

def test_missing_centroids_raises(
self, simple_conv_linear_model, simple_model_input, basic_config
):
"""Test that a state dict with missing centroids raises an error."""
src = KMeansPalettizer(copy.deepcopy(simple_conv_linear_model), basic_config)
state_dict = src.prepare((simple_model_input,)).state_dict()
pruned = {k: v for k, v in state_dict.items() if not k.endswith(".centroids")}

dst = KMeansPalettizer(copy.deepcopy(simple_conv_linear_model), basic_config)
with pytest.raises(RuntimeError, match="does not match the palettizer config"):
dst.prepare((simple_model_input,), state_dict=pruned)
Loading