From c8cd78b1a486eb41d3b73b87c28ac137217fbb7d Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Sat, 13 Jun 2026 14:33:16 -0600 Subject: [PATCH] Allow nested transformer block names for quantization, lora targeting, quantizing --- .../diffusion_models/hidream/hidream_o1_model.py | 3 ++- jobs/process/BaseSDTrainProcess.py | 8 +++++++- toolkit/lora_special.py | 5 ++++- toolkit/util/quantize.py | 8 +++++++- 4 files changed, 20 insertions(+), 4 deletions(-) diff --git a/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py b/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py index 41bf5e55..80790877 100644 --- a/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py +++ b/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py @@ -127,6 +127,7 @@ class HidreamO1Model(BaseModel): super().__init__( device, model_config, dtype, custom_pipeline, noise_scheduler, **kwargs ) + self.use_old_lokr_format = False self.is_flow_matching = True self.is_transformer = True self.target_lora_modules = ["Qwen3VLForConditionalGeneration"] @@ -524,7 +525,7 @@ class HidreamO1Model(BaseModel): return self.arch def get_transformer_block_names(self) -> Optional[List[str]]: - return ["layers"] + return ["model.language_model.layers"] def convert_lora_weights_before_save(self, state_dict): new_sd = {} diff --git a/jobs/process/BaseSDTrainProcess.py b/jobs/process/BaseSDTrainProcess.py index 0ae98840..bfd35365 100644 --- a/jobs/process/BaseSDTrainProcess.py +++ b/jobs/process/BaseSDTrainProcess.py @@ -2135,7 +2135,13 @@ class BaseSDTrainProcess(BaseTrainProcess): compiled_block_count = 0 for attr_name in BLOCK_LIST_ATTRS: - block_list = getattr(inner_unet, attr_name, None) + # attr_name may be a dotted path for models that nest their + # blocks (e.g. hidream_o1's "model.language_model.layers"). + block_list = inner_unet + for part in attr_name.split('.'): + block_list = getattr(block_list, part, None) + if block_list is None: + break if block_list is None: continue diff --git a/toolkit/lora_special.py b/toolkit/lora_special.py index 5cb19229..acb1e870 100644 --- a/toolkit/lora_special.py +++ b/toolkit/lora_special.py @@ -370,7 +370,10 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork): transformer_block_names = base_model.get_transformer_block_names() if transformer_block_names is not None: - if not any([name in lora_name for name in transformer_block_names]): + # match against clean_name (dotted) so block names can be + # dotted paths (e.g. "model.language_model.layers"); lora_name + # has dots replaced with "$$"/"_" and wouldn't match. + if not any([block_name in clean_name for block_name in transformer_block_names]): skip = True else: if self.is_pixart: diff --git a/toolkit/util/quantize.py b/toolkit/util/quantize.py index 03c4d78a..3b0af71d 100644 --- a/toolkit/util/quantize.py +++ b/toolkit/util/quantize.py @@ -297,7 +297,13 @@ def quantize_model( all_blocks: List[torch.nn.Module] = [] transformer_block_names = base_model.get_transformer_block_names() for name in transformer_block_names: - block_list = getattr(model_to_quantize, name, None) + # name may be a dotted path for models that nest their blocks + # (e.g. hidream_o1's "model.language_model.layers"). + block_list = model_to_quantize + for part in name.split('.'): + block_list = getattr(block_list, part, None) + if block_list is None: + break if block_list is not None: all_blocks += list(block_list) base_model.print_and_status_update(