Repository navigation
Allow palettizer prepare to skip centroid calculation given state dict - #120
Conversation
| """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() |
There was a problem hiding this comment.
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
left a comment
There was a problem hiding this comment.
Looks good, thanks for this change!
| if self._indices_stale: | ||
| destination.pop(prefix + "indices", None) |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
_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.
| if not info.fp_module.check_compatible(info.weight): | ||
| info.fp_module._disabled = True |
There was a problem hiding this comment.
Minor readability suggestion:
| 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( |
There was a problem hiding this comment.
Can we also point to what we were trying to do?
Something like "unable to load state dict"
cda9bda to
8b6ff88
Compare
This PR enables KMeansPalettizer to skip calculating centroids if an optional
state_dictargument is passed toprepare().Prior to this PR,
prepare()always did the following:self._handler.prepare()to walk the model and insert palettizer modulesdisabledWith these changes, when
state_dictis not provided,prepare()'s behavior is unchanged from above. Whenstate_dictis present, the following is done instead:self._handler.prepare()to walk the model and insert palettizer modules (unchanged)check_compatibleto see if it has shapes/dims incompatible with palettization; if so, mark as 'disabled`.disabledpalettizersAdditional changes:
reshape_and_blockas a separate subutility inscale_reshape_and_blockto allow for use incheck_compatible. Meta tensors are used incheck_compatibleto keep the runtime cheap and quick.save_to_state_dictlogic added to excludeindicesfrom being saved tostate_dictif the_indices_staleflag 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.