MCPcopy Create free account
hub / github.com/Vahe1994/AQLM / load_dequantized_model

Function load_dequantized_model

src/modelutils.py:235–248  ·  view source on GitHub ↗

Load quantized model by dequantizing it

(model, load_path)

Source from the content-addressed store, hash-verified

233
234
235def load_dequantized_model(model, load_path):
236 """Load quantized model by dequantizing it"""
237 layers = get_layers(model)
238 for layer_index in range(len(layers)):
239 print("layer", layer_index)
240 layer = layers[layer_index]
241 quant_layer = torch.load(os.path.join(load_path, str(layer_index) + ".pth"), map_location="cpu")
242 for module in quant_layer.modules():
243 if isinstance(module, QuantizedWeight):
244 if not hasattr(module, "codes_storage"):
245 module.codes_storage = None # backwards compatibility
246 layers[layer_index] = load_linear_layers(layer, quant_layer, model)
247 model.load_state_dict(torch.load(os.path.join(load_path, "not_quantized_weights.pt")), strict=False)
248 return model
249
250
251def load_quantized_model(model, load_path):

Callers 1

from_pretrained_aqlmFunction · 0.90

Calls 3

get_layersFunction · 0.85
load_linear_layersFunction · 0.85
load_state_dictMethod · 0.80

Tested by

no test coverage detected