mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
refactor(ideogram4): register pipeline and load via from_pretrained
Register the subclass with generic.set_pipeline and build it with from_pretrained, passing the SDNQ transformers and shared text encoder while the vae, scheduler, and tokenizer load from the repo.
This commit is contained in:
@@ -40,7 +40,6 @@ pipelines = {
|
||||
'HiDream': getattr(diffusers, 'HiDreamImagePipeline', None),
|
||||
'HunyuanDiT': getattr(diffusers, 'HunyuanDiTPipeline', None),
|
||||
'HunyuanImage': getattr(diffusers, 'HunyuanImagePipeline', None),
|
||||
'Ideogram4': getattr(diffusers, 'Ideogram4Pipeline', None),
|
||||
'JoyEdit': getattr(diffusers, 'JoyImageEditPipeline', None),
|
||||
'Kandinsky21': getattr(diffusers, 'KandinskyCombinedPipeline', None),
|
||||
'Kandinsky22': getattr(diffusers, 'KandinskyV22CombinedPipeline', None),
|
||||
@@ -70,6 +69,7 @@ pipelines = {
|
||||
'FLEX': None,
|
||||
'HiDreamO1': None,
|
||||
'HunyuanImage3': None,
|
||||
'Ideogram4': None,
|
||||
'Lens': None,
|
||||
'LuminaDiMOO': None,
|
||||
'Meissonic': None,
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
import json
|
||||
import diffusers
|
||||
from transformers import AutoTokenizer
|
||||
from transformers.models.qwen3_vl import Qwen3VLModel
|
||||
from modules import shared, devices, sd_models
|
||||
from modules import shared, devices, sd_models, model_quant
|
||||
from modules.logger import log
|
||||
from pipelines import generic
|
||||
|
||||
@@ -70,8 +69,10 @@ def load_ideogram4(checkpoint_info, diffusers_load_config=None):
|
||||
diffusers_load_config = {}
|
||||
repo_id = sd_models.path_to_repo(checkpoint_info)
|
||||
sd_models.hf_auth_check(checkpoint_info)
|
||||
log.debug(f'Load model: type=Ideogram4 repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}')
|
||||
load_args, _ = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
|
||||
log.debug(f'Load model: type=Ideogram4 repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
|
||||
|
||||
generic.set_pipeline('Ideogram4', Ideogram4Pipeline)
|
||||
if repo_id is None or repo_id.lower() == 'none':
|
||||
return None
|
||||
|
||||
@@ -81,19 +82,17 @@ def load_ideogram4(checkpoint_info, diffusers_load_config=None):
|
||||
unconditional_transformer = generic.load_transformer(repo_id, cls_name=cls, subfolder="unconditional_transformer", load_config=diffusers_load_config)
|
||||
pin_transformers_if_fit(transformer, unconditional_transformer)
|
||||
# shared_te_map redirects to the shared Qwen3-VL repo (deduped with VQA + prompt-enhance);
|
||||
# the bundled text_encoder is the fallback when sharing is off.
|
||||
# the bundled text_encoder is the fallback when sharing is off. The vae, tokenizer, and
|
||||
# scheduler load from the repo via from_pretrained.
|
||||
text_encoder = generic.load_text_encoder(repo_id, cls_name=Qwen3VLModel, load_config=diffusers_load_config)
|
||||
tokenizer = AutoTokenizer.from_pretrained(repo_id, subfolder="tokenizer", cache_dir=shared.opts.diffusers_dir)
|
||||
vae = diffusers.AutoencoderKLFlux2.from_pretrained(repo_id, subfolder="vae", cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype)
|
||||
scheduler = diffusers.FlowMatchEulerDiscreteScheduler.from_pretrained(repo_id, subfolder="scheduler", cache_dir=shared.opts.diffusers_dir)
|
||||
|
||||
pipe = Ideogram4Pipeline(
|
||||
scheduler=scheduler,
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
pipe = Ideogram4Pipeline.from_pretrained(
|
||||
repo_id,
|
||||
cache_dir=shared.opts.diffusers_dir,
|
||||
transformer=transformer,
|
||||
unconditional_transformer=unconditional_transformer,
|
||||
text_encoder=text_encoder,
|
||||
**load_args,
|
||||
)
|
||||
# The pipeline decodes internally; the CFG scale slider drives guidance_scale, which is
|
||||
# mutually exclusive with the pipeline's default per-step guidance_schedule.
|
||||
@@ -101,6 +100,6 @@ def load_ideogram4(checkpoint_info, diffusers_load_config=None):
|
||||
# JSON captions must pass through verbatim; skip styles/wildcards that would strip the braces.
|
||||
pipe.keep_prompts = True # pylint: disable=attribute-defined-outside-init
|
||||
|
||||
del transformer, unconditional_transformer, text_encoder, vae
|
||||
del transformer, unconditional_transformer, text_encoder
|
||||
devices.torch_gc(force=True, reason='load')
|
||||
return pipe
|
||||
|
||||
Reference in New Issue
Block a user