Skip to content

Commit f47e930

Browse files
authored
Support quantization with disk offloading (#1554)
* fix quant disk offload * update text_encoder quant
1 parent ab12bf4 commit f47e930

3 files changed

Lines changed: 75 additions & 28 deletions

File tree

‎diffsynth/configs/model_configs.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1372,11 +1372,11 @@
13721372
},
13731373
{
13741374
# Example: ModelConfig(model_id="DiffSynth-Studio/MiniMax-H3-NF4", origin_file_pattern="minimax-h3-text-encoder-nf4.safetensors")
1375-
"model_hash": "14c8a9ac1e38161b6989158689d8b28b",
1375+
"model_hash": "297933c3a2b0fc4d4dfee30e34c566b8",
13761376
"model_name": "minimax_h3_text_encoder",
13771377
"model_class": "diffsynth.models.minimax_h3_text_encoder.MiniMaxH3TextEncoder",
13781378
"state_dict_converter": "diffsynth.utils.state_dict_converters.minimax_h3_text_encoder.MiniMaxH3TextEncoderStateDictConverter",
1379-
"quant_config": {"method": "bitsandbytes_nf4", "load_prequantized": True},
1379+
"quant_config": {"method": "bitsandbytes_nf4", "load_prequantized": True, "exclude_modules": ["qkv", "proj", "linear_fc1", "linear_fc2"]},
13801380
},
13811381
{
13821382
# Example: ModelConfig(model_id="MiniMax/MiniMax-H3", origin_file_pattern="FL2VA/transformer/model*.safetensors")

‎diffsynth/core/loader/model.py‎

Lines changed: 27 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -32,29 +32,37 @@ def load_model(model_class, path, config=None, torch_dtype=torch.bfloat16, devic
3232
model = enable_vram_management(model, module_map, vram_config=vram_config, disk_map=disk_map, vram_limit=vram_limit)
3333
elif quantize is not None and module_map is not None:
3434
if "disk" in vram_config.values():
35-
raise ValueError("Model quantization is incompatible with disk offload.")
36-
offload_device = vram_config["offload_device"]
37-
computation_device = vram_config["computation_device"]
38-
computation_dtype = vram_config["computation_dtype"]
39-
offload_dtype = vram_config["offload_dtype"]
40-
load_dtype = None if quantize.load_prequantized else computation_dtype
41-
if state_dict is None: state_dict = DiskMap(path, offload_device, torch_dtype=load_dtype)
42-
if state_dict_converter is not None:
43-
state_dict = state_dict_converter(state_dict)
35+
if not quantize.load_prequantized:
36+
raise ValueError("Disk offload with quantization is only supported for pre-quantized checkpoints (load_prequantized=True).")
37+
devices = [vram_config[k] for k in ("offload_device", "onload_device", "preparing_device", "computation_device")]
38+
load_device = [d for d in devices if d != "disk"][0]
39+
disk_map = DiskMap(path, load_device, torch_dtype=None, state_dict_converter=state_dict_converter)
40+
metadata = load_metadata_from_safetensors(path[0] if isinstance(path, list) else path)
41+
model = quantize.prepare_for_prequantized_load(model, compute_dtype=vram_config["computation_dtype"])
42+
model = enable_vram_management(model, module_map, vram_config=vram_config, disk_map=disk_map, vram_limit=vram_limit, quantize=quantize, metadata=metadata)
4443
else:
45-
state_dict = {i: state_dict[i] for i in state_dict}
44+
offload_device = vram_config["offload_device"]
45+
computation_device = vram_config["computation_device"]
46+
computation_dtype = vram_config["computation_dtype"]
47+
offload_dtype = vram_config["offload_dtype"]
48+
load_dtype = None if quantize.load_prequantized else computation_dtype
49+
if state_dict is None: state_dict = DiskMap(path, offload_device, torch_dtype=load_dtype)
50+
if state_dict_converter is not None:
51+
state_dict = state_dict_converter(state_dict)
52+
else:
53+
state_dict = {i: state_dict[i] for i in state_dict}
4654

47-
if quantize.load_prequantized:
48-
model = quantize.prepare_for_prequantized_load(model, compute_dtype=computation_dtype)
49-
state_dict = quantize.unflatten_state_dict(state_dict, load_metadata_from_safetensors(path))
55+
if quantize.load_prequantized:
56+
model = quantize.prepare_for_prequantized_load(model, compute_dtype=computation_dtype)
57+
state_dict = quantize.unflatten_state_dict(state_dict, load_metadata_from_safetensors(path))
5058

51-
model.load_state_dict(state_dict, assign=True)
52-
state_dict = None
59+
model.load_state_dict(state_dict, assign=True)
60+
state_dict = None
5361

54-
model = quantize.quantize_model(model, compute_device=computation_device, model_device=offload_device)
55-
model = quantize.dequantize_model(model, compute_dtype=computation_dtype, compute_device=computation_device, model_device=offload_device)
56-
model = model.to(dtype=offload_dtype, device=offload_device)
57-
model = enable_vram_management(model, module_map, vram_config=vram_config, disk_map=None, vram_limit=vram_limit, quantize=quantize)
62+
model = quantize.quantize_model(model, compute_device=computation_device, model_device=offload_device)
63+
model = quantize.dequantize_model(model, compute_dtype=computation_dtype, compute_device=computation_device, model_device=offload_device)
64+
model = model.to(dtype=offload_dtype, device=offload_device)
65+
model = enable_vram_management(model, module_map, vram_config=vram_config, disk_map=None, vram_limit=vram_limit, quantize=quantize)
5866
elif quantize is not None:
5967
# Weight-only quantization (see `diffsynth.core.quant`), isolated from the normal path below.
6068
if quantize.load_prequantized:

‎diffsynth/core/vram/layers.py‎

Lines changed: 46 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -451,12 +451,15 @@ def __init__(
451451
vram_limit: float = None,
452452
name: str = "",
453453
disk_map: DiskMap = None,
454+
quantize=None,
455+
metadata: dict = None,
454456
**kwargs
455457
):
456-
if "disk" in (offload_dtype, offload_device, onload_dtype, onload_device):
458+
disk_offload = "disk" in (offload_dtype, offload_device, onload_dtype, onload_device)
459+
if disk_offload and (disk_map is None or quantize is None):
457460
raise ValueError(
458-
"Layer quantization is incompatible with disk offload: DiskMap holds plain "
459-
"tensors only and cannot rebuild packed weights with their quant state."
461+
"Disk offload for quantized layers requires both `disk_map` and `quantize`, "
462+
"so each layer can rebuild its packed weight and quant state lazily."
460463
)
461464
super().__init__(
462465
offload_dtype,
@@ -471,8 +474,32 @@ def __init__(
471474
)
472475
self.module = module
473476
self.name = name
477+
self.disk_offload = disk_offload
478+
self.disk_map = disk_map
479+
self.quantize = quantize
480+
self.metadata = metadata
481+
self._required_keys = None
474482
self.init_lora_hotload()
475483

484+
def _disk_required_keys(self):
485+
if self._required_keys is None:
486+
weight_prefix = self.name + ".weight."
487+
self._required_keys = [
488+
key for key in self.disk_map
489+
if key == self.name + ".weight" or key.startswith(weight_prefix) or key == self.name + ".bias"
490+
]
491+
return self._required_keys
492+
493+
def _load_from_disk(self, device, target=None):
494+
module = self.module if target is None else target
495+
prefix = self.name + "."
496+
subdict = {key: self.disk_map[key] for key in self._disk_required_keys()}
497+
rebuilt = self.quantize.unflatten_state_dict(subdict, self.metadata or {})
498+
state = {key[len(prefix):]: value for key, value in rebuilt.items()}
499+
module.load_state_dict(state, assign=True)
500+
module.to(device=device)
501+
return module
502+
476503
def _module_device(self):
477504
tensor = next(self.module.parameters(), None)
478505
if tensor is None:
@@ -481,23 +508,35 @@ def _module_device(self):
481508

482509
def offload(self):
483510
if self.state != 0:
484-
self.module.to(device=self.offload_device)
511+
if self.disk_offload:
512+
self.module = self.quantize.backend.create_quantized_linear_shell(self.module, self.computation_dtype)
513+
else:
514+
self.module.to(device=self.offload_device)
485515
self.state = 0
486516

487517
def onload(self):
488518
if self.state < 1:
489-
self.module.to(device=self.onload_device)
519+
if self.disk_offload and self.onload_device != "disk" and self.offload_device == "disk":
520+
self._load_from_disk(self.onload_device)
521+
elif self.onload_device != "disk":
522+
self.module.to(device=self.onload_device)
490523
self.state = 1
491524

492525
def preparing(self):
493526
if self.state != 2:
494-
self.module.to(device=self.preparing_device)
527+
if self.disk_offload and self.preparing_device != "disk" and self.onload_device == "disk":
528+
self._load_from_disk(self.preparing_device)
529+
elif self.preparing_device != "disk":
530+
self.module.to(device=self.preparing_device)
495531
self.state = 2
496532

497533
def computation_module(self):
498534
device = self.preparing_device if self.state == 2 else self.onload_device
499535
if device == self.computation_device:
500536
return self.module
537+
if self.disk_offload and device == "disk":
538+
transient = self.quantize.backend.create_quantized_linear_shell(self.module, self.computation_dtype)
539+
return self._load_from_disk(self.computation_device, target=transient)
501540
return copy.deepcopy(self.module).to(device=self.computation_device)
502541

503542
def forward(self, x, *args, **kwargs):
@@ -522,7 +561,7 @@ def enable_vram_management_recursively(model: torch.nn.Module, module_map: dict,
522561
for name, module in model.named_children():
523562
layer_name = name if name_prefix == "" else name_prefix + "." + name
524563
if quantize is not None and quantize.is_quantized_linear(module):
525-
module_ = AutoWrappedQuantizedModule(module, **vram_config, vram_limit=vram_limit, name=layer_name, disk_map=disk_map, **kwargs)
564+
module_ = AutoWrappedQuantizedModule(module, **vram_config, vram_limit=vram_limit, name=layer_name, disk_map=disk_map, quantize=quantize, **kwargs)
526565
setattr(model, name, module_)
527566
continue
528567
for source_module, target_module in module_map.items():

0 commit comments

Comments
 (0)