Skip to content

Allow palettizer prepare to skip centroid calculation given state dict - #120

Merged
crowbat merged 4 commits into
apple:mainfrom
crowbat:u/k_hsieh/palettizer_prepare_skip_centroid_calc
Oct 1, 2026
Merged

crowbat merged 4 commits into
apple:mainfrom
crowbat:u/k_hsieh/palettizer_prepare_skip_centroid_calc

Conversation

@crowbat

@crowbat crowbat commented Sep 25, 2026

Copy link
Copy Markdown
Contributor

This PR enables KMeansPalettizer to skip calculating centroids if an optional state_dict argument is passed to prepare().

Prior to this PR, prepare() always did the following:

  1. Call self._handler.prepare() to walk the model and insert palettizer modules
  2. Compute centroids and indices by calling forward passes for each palettized module
    • As part of calling forward passes, modules with incompatible shapes/dims for palettization will be marked as disabled
  3. Remove all palettized modules which are disabled

With these changes, when state_dict is not provided, prepare()'s behavior is unchanged from above. When state_dict is present, the following is done instead:

  1. Call self._handler.prepare() to walk the model and insert palettizer modules (unchanged)
  2. For every palettized module, call check_compatible to see if it has shapes/dims incompatible with palettization; if so, mark as 'disabled`.
  3. Remove all disabled palettizers
  4. Load the state dict strictly to catch any state dicts with more keys than expected. Then check all palettized modules to make sure all centroids are present, to catch state dicts with less keys than expected.

Additional changes:

  • Extraction of reshape_and_block as a separate subutility in scale_reshape_and_block to allow for use in check_compatible. Meta tensors are used in check_compatible to keep the runtime cheap and quick.
  • Custom save_to_state_dict logic added to exclude indices from being saved to state_dict if the _indices_stale flag is set to True. When loading a state dict without indices, this results in centroids being loaded in (allowing for the skipping of centroid calculation during prepare), however indices will will need to be assigned on the next forward pass of the model.

"""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 = prepared.state_dict()

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.

Just for completeness, could we have one of the tests, also save the state_dict into a .pt file (in a temp folder), and then load it from the pt file in the next prepare

@u-simha u-simha left a comment

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.

Looks good, thanks for this change!

Comment on lines +257 to +258
if self._indices_stale:
destination.pop(prefix + "indices", None)

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.

Comment on lines +257 to +258
if not info.fp_module.check_compatible(info.weight):
info.fp_module._disabled = True

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.

Minor readability suggestion:

Suggested change
if not info.fp_module.check_compatible(info.weight):
info.fp_module._disabled = True
compatible = info.fp_module.check_compatible(info.weight)
info.fp_module._disabled = not compatible

for info in self._collect_fake_palett_info(to_cpu=False):
fp = info.fp_module
if fp.centroids is None:
raise RuntimeError(

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.

Can we also point to what we were trying to do?

Something like "unable to load state dict"

@crowbat
crowbat force-pushed the u/k_hsieh/palettizer_prepare_skip_centroid_calc branch from cda9bda to 8b6ff88 Compare October 1, 2026 16:34
@crowbat
crowbat enabled auto-merge (squash) October 1, 2026 16:40
@crowbat
crowbat merged commit 1806023 into apple:main Oct 1, 2026
14 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants