import sdnq on video prequant model load

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-11-01 11:20:30 -04:00
parent 372770b285
commit c1d87a14eb
6 changed files with 17 additions and 6 deletions
+4
View File
@@ -5,7 +5,11 @@
- add inline wildcards using curly braces syntax
- fix inpaint
- fix model type detection
- fix version detection when cloned with `.git` suffix
- guard against multi-controlnet in hires
- fix full-screen image viewer buttons with non-standard ui theme
- init `sdnq` on video model load
- control tab show override section
## Update for 2025-10-31
+2 -1
View File
@@ -351,7 +351,8 @@ def create_refresh_button(refresh_component, refresh_method, refreshed_args = No
def create_override_inputs(tab): # pylint: disable=unused-argument
with gr.Row(elem_id=f"{tab}_override_settings_row"):
override_settings = gr.Dropdown([], value=None, label="Override settings", visible=False, elem_id=f"{tab}_override_settings", multiselect=True)
visible = tab == 'control'
override_settings = gr.Dropdown([], value=None, label="Override settings", visible=visible, elem_id=f"{tab}_override_settings", multiselect=True)
override_settings.change(fn=lambda x: gr.Dropdown.update(visible=len(x) > 0), inputs=[override_settings], outputs=[override_settings])
return override_settings
+6
View File
@@ -58,6 +58,9 @@ def load_model(selected: models_def.Model):
selected.te_folder = 'text_encoder'
selected.te_revision = None
if 'sdnq-' in (selected.te or selected.repo).lower():
from modules import sdnq # pylint: disable=unused-import # register to diffusers and transformers
shared.log.debug(f'Video load: module=te repo="{selected.te or selected.repo}" folder="{selected.te_folder}" cls={selected.te_cls.__name__} quant={model_quant.get_quant_type(quant_args)}')
kwargs["text_encoder"] = selected.te_cls.from_pretrained(
pretrained_model_name_or_path=selected.te or selected.repo,
@@ -91,6 +94,9 @@ def load_model(selected: models_def.Model):
else:
shared.log.debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" module="{dit_folder}" folder="{dit_folder}" cls={selected.dit_cls.__name__} skip')
if 'sdnq-' in (selected.dit or selected.repo).lower():
from modules import sdnq # pylint: disable=unused-import # register to diffusers and transformers
if selected.dit_folder is None:
selected.dit_folder = ['transformer']
if isinstance(selected.dit_folder, list) or isinstance(selected.dit_folder, tuple):
+2 -2
View File
@@ -24,7 +24,7 @@ from transformers import AutoTokenizer, CLIPImageProcessor, CLIPVisionModel, UMT
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.image_processor import PipelineImageInput
from diffusers.loaders import WanLoraLoaderMixin
from diffusers.models import AutoencoderKLWan, WanTransformer3DModel
from diffusers.models import AutoencoderKLWan, WanTransformer3DModel # pylint: disable=unused-import # register to diffusers
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import is_ftfy_available, is_torch_xla_available, logging, replace_example_docstring
from diffusers.utils.torch_utils import randn_tensor
@@ -731,7 +731,7 @@ class ChronoEditPipeline(DiffusionPipeline, WanLoraLoaderMixin):
self._current_timestep = None
if not output_type == "latent":
if output_type != "latent":
latents = latents.to(self.vae.dtype)
latents_mean = (
torch.tensor(self.vae.config.latents_mean)
+2 -2
View File
@@ -188,7 +188,7 @@ class ChronoEditRotaryPosEmbed(nn.Module):
self.freqs = torch.cat(freqs, dim=1)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
batch_size, _num_channels, num_frames, height, width = hidden_states.shape
_batch_size, _num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w
@@ -419,7 +419,7 @@ class ChronoEditTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, Fr
)
batch_size, _num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size
p_t, p_h, p_w = self.config.patch_size # pylint: disable=no-member
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w