mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
minor update
This commit is contained in:
@@ -5,10 +5,11 @@ import diffusers
|
||||
import transformers
|
||||
from safetensors.torch import load_file
|
||||
from huggingface_hub import hf_hub_download
|
||||
from modules import shared, devices
|
||||
from modules import shared, devices, modelloader
|
||||
|
||||
|
||||
debug = shared.log.trace if os.environ.get('SD_LOAD_DEBUG', None) is not None else lambda *args, **kwargs: None
|
||||
base_repo = 'black-forest-labs/FLUX.1-dev'
|
||||
|
||||
|
||||
def get_quant(file_path):
|
||||
@@ -132,6 +133,7 @@ def load_transformer(file_path, transformer): # triggered by opts.sd_unet change
|
||||
def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_checkpoint change
|
||||
quant = get_quant(checkpoint_info.path)
|
||||
shared.log.debug(f'Loading FLUX: model="{checkpoint_info.name}" unet="{shared.opts.sd_unet}" t5="{shared.opts.sd_text_encoder}" vae="{shared.opts.sd_vae}" quant={quant} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}')
|
||||
modelloader.hf_login()
|
||||
|
||||
transformer = None
|
||||
text_encoder_2 = None
|
||||
@@ -171,5 +173,5 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
|
||||
if vae is not None:
|
||||
components['vae'] = vae
|
||||
debug(f'Loading FLUX: preloaded={list(components)}')
|
||||
pipe = diffusers.FluxPipeline.from_pretrained('black-forest-labs/FLUX.1-dev', cache_dir=shared.opts.diffusers_dir, **components, **diffusers_load_config)
|
||||
pipe = diffusers.FluxPipeline.from_pretrained(base_repo, cache_dir=shared.opts.diffusers_dir, **components, **diffusers_load_config)
|
||||
return pipe
|
||||
|
||||
@@ -190,7 +190,7 @@ def load_flux_nf4(checkpoint_info, diffusers_load_config, transformer, text_enco
|
||||
|
||||
with init_empty_weights():
|
||||
from diffusers import FluxTransformer2DModel
|
||||
config = FluxTransformer2DModel.load_config("black-forest-labs/FLUX.1-dev", subfolder="transformer")
|
||||
config = FluxTransformer2DModel.load_config(os.path.join('configs', 'flux'), subfolder="transformer")
|
||||
transformer = FluxTransformer2DModel.from_config(config).to(devices.dtype)
|
||||
expected_state_dict_keys = list(transformer.state_dict().keys())
|
||||
|
||||
|
||||
Reference in New Issue
Block a user