Load quantized model by dequantizing it
(model, load_path)
| 233 | |
| 234 | |
| 235 | def 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 | |
| 251 | def load_quantized_model(model, load_path): |
no test coverage detected