Add auto detection based on model keys for v_pred and add auto detection for nyaflow

This commit is contained in:
Disty0
2025-08-25 18:32:45 +03:00
parent b1ad9fb650
commit 073d6cec77
+21 -3
View File
@@ -539,16 +539,34 @@ def load_diffuser_file(model_type, pipeline, checkpoint_info, diffusers_load_con
def set_overrides(sd_model, checkpoint_info):
if 'bigaspv25' in checkpoint_info.name.lower():
checkpoint_info_name = checkpoint_info.name.lower()
if 'bigaspv25' in checkpoint_info_name or 'nyaflow' in checkpoint_info_name:
scheduler_config = sd_model.scheduler.config
scheduler_config['prediction_type'] = 'flow_prediction'
scheduler_config['use_flow_sigmas'] = True
scheduler_config['beta_schedule'] = 'linear'
sd_model.scheduler = diffusers.UniPCMultistepScheduler.from_config(scheduler_config)
shared.log.info(f'Setting override: model="{checkpoint_info.name}" component=scheduler prediction="flow-prediction"')
if 'vpred' in checkpoint_info.name.lower() or 'v-pred' in checkpoint_info.name.lower():
elif 'vpred' in checkpoint_info_name or 'v-pred' in checkpoint_info_name or 'v_pred' in checkpoint_info_name:
scheduler_config = sd_model.scheduler.config
scheduler_config['prediction_type'] = 'v_prediction'
scheduler_config['rescale_betas_zero_snr'] = True
sd_model.scheduler = diffusers.EulerDiscreteScheduler.from_config(scheduler_config)
shared.log.info(f'Setting override: model="{checkpoint_info.name}" component=scheduler prediction="v-prediction"')
shared.log.info(f'Setting override: model="{checkpoint_info.name}" component=scheduler prediction="v-prediction" rescale=True')
elif checkpoint_info.path.lower().endswith('.safetensors'):
try:
from safetensors import safe_open
with safe_open(checkpoint_info.path, framework='pt') as f:
keys = f.keys()
if 'v_pred' in keys: # NoobAI VPred models added empty v_pred and ztsnr keys
scheduler_config = sd_model.scheduler.config
scheduler_config['prediction_type'] = 'v_prediction'
if 'ztsnr' in keys:
scheduler_config['rescale_betas_zero_snr'] = True
sd_model.scheduler = diffusers.EulerDiscreteScheduler.from_config(scheduler_config)
shared.log.info(f'Setting override: model="{checkpoint_info.name}" component=scheduler prediction="v-prediction" rescale={scheduler_config.get("rescale_betas_zero_snr", False)}')
except Exception as e:
shared.log.debug(f'Setting override from keys failed: {e}')
def set_defaults(sd_model, checkpoint_info):