close
Skip to content

Commit fed7b18

Browse files
authored
support dynamic quant for LoRA training (#1622)
* support dynamic quant for training * resolve registered config * update default is_differentiable to True * move quant to special training * remove reduant
1 parent db5b335 commit fed7b18

65 files changed

Lines changed: 238 additions & 31 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎diffsynth/core/quant/base.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,7 @@ def __init__(self, config=None):
5252
def capabilities(self) -> dict:
5353
return {
5454
"is_serializable": False,
55-
"is_differentiable": False,
55+
"is_differentiable": True,
5656
"is_compileable": False,
5757
"requires_calibration": False,
5858
}

‎diffsynth/diffusion/parsers.py‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,10 @@ def add_model_config(parser: argparse.ArgumentParser):
3131
parser.add_argument("--resume_from_checkpoint", default=None, type=str, help="Resume training from checkpoint file. Only single model training is supported.")
3232
return parser
3333

34+
def add_quant_config(parser: argparse.ArgumentParser):
35+
parser.add_argument("--quant_options", type=str, default=None, help="Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. bitsandbytes_nf4), and `exclude_modules` optionally lists layers kept in full precision.")
36+
return parser
37+
3438
def add_training_config(parser: argparse.ArgumentParser):
3539
parser.add_argument("--learning_rate", type=float, default=1e-4, help="Learning rate.")
3640
parser.add_argument("--num_epochs", type=int, default=1, help="Number of epochs.")
@@ -106,6 +110,7 @@ def add_dmd2_config(parser: argparse.ArgumentParser):
106110
def add_general_config(parser: argparse.ArgumentParser):
107111
parser = add_dataset_base_config(parser)
108112
parser = add_model_config(parser)
113+
parser = add_quant_config(parser)
109114
parser = add_training_config(parser)
110115
parser = add_output_config(parser)
111116
parser = add_lora_config(parser)

‎diffsynth/diffusion/training_module.py‎

Lines changed: 25 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
import torch, json, os, inspect
2-
from ..core import ModelConfig, load_state_dict
2+
from ..core import ModelConfig, load_state_dict, QuantizeConfig
33
from ..utils.controlnet import ControlNetInput
44
from .base_pipeline import PipelineUnit
55
from peft import LoraConfig, inject_adapter_in_model
@@ -176,9 +176,30 @@ def parse_vram_config(self, fp8=False, offload=False, device="cpu"):
176176
else:
177177
return {}
178178

179-
def parse_model_configs(self, model_paths, model_id_with_origin_paths, fp8_models=None, offload_models=None, device="cpu"):
179+
def parse_quant_options(self, quant_options):
180+
quant_map = {}
181+
if quant_options is None:
182+
return quant_map
183+
for entry in quant_options.split(";"):
184+
if entry == "":
185+
continue
186+
model_string, sep, spec = entry.rpartition(":")
187+
if sep == "":
188+
raise ValueError(f"Failed to parse quant option: `{entry}`. Expected `<model_string>:<method>[/<exclude_modules>]`.")
189+
if spec == "":
190+
continue
191+
method, _, excludes = spec.partition("/")
192+
exclude_modules = excludes.split(",") if excludes != "" else None
193+
quant_config = QuantizeConfig(method=method, exclude_modules=exclude_modules)
194+
if not quant_config.backend.capabilities().get("is_differentiable", True):
195+
raise ValueError(f"Quantization method `{method}` is not differentiable, so it cannot be used for training (frozen quantized layers must pass gradients through to LoRA branches). Choose a method whose backend declares `is_differentiable=True`.")
196+
quant_map[model_string] = quant_config
197+
return quant_map
198+
199+
def parse_model_configs(self, model_paths, model_id_with_origin_paths, fp8_models=None, offload_models=None, quant_options=None, device="cpu"):
180200
fp8_models = [] if fp8_models is None else fp8_models.split(",")
181201
offload_models = [] if offload_models is None else offload_models.split(",")
202+
quant_map = self.parse_quant_options(quant_options)
182203
model_configs = []
183204
if model_paths is not None:
184205
model_paths = json.loads(model_paths)
@@ -188,7 +209,7 @@ def parse_model_configs(self, model_paths, model_id_with_origin_paths, fp8_model
188209
offload=path in offload_models,
189210
device=device
190211
)
191-
model_configs.append(ModelConfig(path=path, **vram_config))
212+
model_configs.append(ModelConfig(path=path, quantize=quant_map.get(path), **vram_config))
192213
if model_id_with_origin_paths is not None:
193214
model_id_with_origin_paths = model_id_with_origin_paths.split(",")
194215
for model_id_with_origin_path in model_id_with_origin_paths:
@@ -198,7 +219,7 @@ def parse_model_configs(self, model_paths, model_id_with_origin_paths, fp8_model
198219
device=device
199220
)
200221
config = self.parse_path_or_model_id(model_id_with_origin_path)
201-
model_configs.append(ModelConfig(model_id=config.model_id, origin_file_pattern=config.origin_file_pattern, **vram_config))
222+
model_configs.append(ModelConfig(model_id=config.model_id, origin_file_pattern=config.origin_file_pattern, quantize=quant_map.get(model_id_with_origin_path), **vram_config))
202223
return model_configs
203224

204225

‎diffsynth/models/model_loader.py‎

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -69,13 +69,12 @@ def default_vram_config(self):
6969
return vram_config
7070

7171
def resolve_quant_config(self, config, quantize):
72-
# No merging: an explicit user config wins wholesale; otherwise the registry
73-
# `quant_config` of a published quantized variant is instantiated as-is.
74-
if quantize is not None:
75-
return quantize
76-
if "quant_config" not in config:
77-
return None
78-
return QuantizeConfig(**config["quant_config"])
72+
registered = config.get("quant_config")
73+
if registered is not None:
74+
if quantize is not None:
75+
print(f"Warning: `{config.get('model_name')}` is already a pre-quantized checkpoint; ignoring the passed quantize option.")
76+
return QuantizeConfig(**registered)
77+
return quantize
7978

8079
def auto_load_model(self, path, vram_config=None, vram_limit=None, clear_parameters=False, state_dict=None, quantize=None):
8180
print(f"Loading models from: {json.dumps(path, indent=4)}")

‎docs/en/Model_Details/ACE-Step.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -124,6 +124,7 @@ Models in the ace_step series are trained uniformly via `examples/ace_step/model
124124
* `--model_id_with_origin_paths`: Model IDs with original paths, separated by commas.
125125
* `--extra_inputs`: Additional input parameters required by the model Pipeline, separated by `,`.
126126
* `--fp8_models`: Models to load in FP8 format, currently only supported for models whose parameters are not updated by gradients.
127+
* `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
127128
* Basic Training Configuration
128129
* `--learning_rate`: Learning rate.
129130
* `--num_epochs`: Number of epochs.

‎docs/en/Model_Details/Anima.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,7 @@ Anima models are trained through [`examples/anima/model_training/train.py`](http
100100
* `--model_id_with_origin_paths`: Model IDs with origin paths (e.g., `"anima-team/anima-1B:text_encoder/*.safetensors"`).
101101
* `--extra_inputs`: Additional pipeline inputs (e.g., `controlnet_inputs` for ControlNet).
102102
* `--fp8_models`: FP8-formatted models (same format as `--model_paths`).
103+
* `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
103104
* Training Configuration
104105
* `--learning_rate`: Learning rate.
105106
* `--num_epochs`: Training epochs.

‎docs/en/Model_Details/Boogu-Image.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,6 +109,7 @@ Models in the boogu_image series are trained uniformly via `examples/boogu_image
109109
* `--model_id_with_origin_paths`: Model IDs with original paths, separated by commas.
110110
* `--extra_inputs`: Additional input parameters required by the model Pipeline, separated by `,`.
111111
* `--fp8_models`: Models to load in FP8 format, currently only supported for models whose parameters are not updated by gradients.
112+
* `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
112113
* Basic Training Configuration
113114
* `--learning_rate`: Learning rate.
114115
* `--num_epochs`: Number of epochs.

‎docs/en/Model_Details/ERNIE-Image.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -96,6 +96,7 @@ ERNIE-Image series models are trained uniformly via [`examples/ernie_image/model
9696
* `--model_id_with_origin_paths`: Model IDs with original paths, e.g., `"PaddlePaddle/ERNIE-Image:transformer/diffusion_pytorch_model*.safetensors"`, separated by commas.
9797
* `--extra_inputs`: Additional input parameters required by the model Pipeline, separated by `,`.
9898
* `--fp8_models`: Models to load in FP8 format, currently only supported for models whose parameters are not updated by gradients.
99+
* `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
99100
* Basic Training Configuration
100101
* `--learning_rate`: Learning rate.
101102
* `--num_epochs`: Number of epochs.

‎docs/en/Model_Details/FLUX.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -175,6 +175,7 @@ FLUX series models are uniformly trained through [`examples/flux/model_training/
175175
* `--model_id_with_origin_paths`: Model IDs with original paths, e.g., `"black-forest-labs/FLUX.1-dev:flux1-dev.safetensors"`. Separated by commas.
176176
* `--extra_inputs`: Extra input parameters required by the model Pipeline, e.g., `controlnet_inputs` when training ControlNet models, separated by `,`.
177177
* `--fp8_models`: Models loaded in FP8 format, consistent with `--model_paths` or `--model_id_with_origin_paths` format. Currently only supports models whose parameters are not updated by gradients (no gradient backpropagation, or gradients only update their LoRA).
178+
* `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
178179
* Training Basic Configuration
179180
* `--learning_rate`: Learning rate.
180181
* `--num_epochs`: Number of epochs.

‎docs/en/Model_Details/FLUX2.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -125,6 +125,7 @@ FLUX.2 series models are uniformly trained through [`examples/flux2/model_traini
125125
* `--model_id_with_origin_paths`: Model IDs with original paths, e.g., `"black-forest-labs/FLUX.2-dev:text_encoder/*.safetensors"`. Separated by commas.
126126
* `--extra_inputs`: Extra input parameters required by the model Pipeline, e.g., `controlnet_inputs` when training ControlNet models, separated by `,`.
127127
* `--fp8_models`: Models loaded in FP8 format, consistent with `--model_paths` or `--model_id_with_origin_paths` format. Currently only supports models whose parameters are not updated by gradients (no gradient backpropagation, or gradients only update their LoRA).
128+
* `--quant_options`: Dynamically quantize loaded models. Semicolon-separated entries, each `<model_string>:<method>[/<exclude_modules>]`, where `<model_string>` matches an entry in `--model_paths`/`--model_id_with_origin_paths`, `method` is a registered method (e.g. `bitsandbytes_nf4`), and `exclude_modules` optionally lists layers kept in full precision.
128129
* Training Basic Configuration
129130
* `--learning_rate`: Learning rate.
130131
* `--num_epochs`: Number of epochs.

0 commit comments

Comments
 (0)