mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
lora-direct with bnb
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
@@ -698,7 +698,7 @@ def control_run(state: str = '',
|
||||
# actual processing
|
||||
if p.is_tile:
|
||||
processed: processing.Processed = tile.run_tiling(p, input_image)
|
||||
if processed is None:
|
||||
if processed is None and p.scripts is not None:
|
||||
processed = p.scripts.run(p, *p.script_args)
|
||||
if processed is None:
|
||||
processed: processing.Processed = processing.process_images(p) # run actual pipeline
|
||||
@@ -706,7 +706,8 @@ def control_run(state: str = '',
|
||||
script_run = True
|
||||
|
||||
# postprocessing
|
||||
processed = p.scripts.after(p, processed, *p.script_args)
|
||||
if p.scripts is not None:
|
||||
processed = p.scripts.after(p, processed, *p.script_args)
|
||||
output = None
|
||||
if processed is not None:
|
||||
output = processed.images
|
||||
|
||||
+26
-12
@@ -314,22 +314,29 @@ def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.n
|
||||
if weights_backup is None and wanted_names != (): # pylint: disable=C1803
|
||||
weight = getattr(self, 'weight', None)
|
||||
self.network_weights_backup = None
|
||||
if shared.opts.lora_fuse_diffusers:
|
||||
self.network_weights_backup = True
|
||||
elif getattr(weight, "quant_type", None) in ['nf4', 'fp4']:
|
||||
if getattr(weight, "quant_type", None) in ['nf4', 'fp4']:
|
||||
if bnb is None:
|
||||
bnb = model_quant.load_bnb('Load network: type=LoRA', silent=True)
|
||||
if bnb is not None:
|
||||
with devices.inference_context():
|
||||
self.network_weights_backup = bnb.functional.dequantize_4bit(weight, quant_state=weight.quant_state, quant_type=weight.quant_type, blocksize=weight.blocksize,)
|
||||
if shared.opts.lora_fuse_diffusers:
|
||||
self.network_weights_backup = True
|
||||
else:
|
||||
self.network_weights_backup = bnb.functional.dequantize_4bit(weight, quant_state=weight.quant_state, quant_type=weight.quant_type, blocksize=weight.blocksize,)
|
||||
self.quant_state = weight.quant_state
|
||||
self.quant_type = weight.quant_type
|
||||
self.blocksize = weight.blocksize
|
||||
else:
|
||||
weights_backup = weight.clone()
|
||||
self.network_weights_backup = weights_backup.to(devices.cpu)
|
||||
if shared.opts.lora_fuse_diffusers:
|
||||
self.network_weights_backup = True
|
||||
else:
|
||||
weights_backup = weight.clone()
|
||||
self.network_weights_backup = weights_backup.to(devices.cpu)
|
||||
else:
|
||||
self.network_weights_backup = weight.clone().to(devices.cpu)
|
||||
if shared.opts.lora_fuse_diffusers:
|
||||
self.network_weights_backup = True
|
||||
else:
|
||||
self.network_weights_backup = weight.clone().to(devices.cpu)
|
||||
|
||||
bias_backup = getattr(self, "network_bias_backup", None)
|
||||
if bias_backup is None:
|
||||
@@ -408,13 +415,20 @@ def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.
|
||||
if updown is not None:
|
||||
if deactivate:
|
||||
updown *= -1
|
||||
try:
|
||||
new_weight = self.weight.to(devices.device) + updown.to(devices.device)
|
||||
except Exception:
|
||||
new_weight = self.weight + updown
|
||||
if getattr(self, "quant_type", None) in ['nf4', 'fp4'] and bnb is not None:
|
||||
self.weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize)
|
||||
try: # TODO lora-direct with bnb
|
||||
weight = bnb.functional.dequantize_4bit(self.weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize)
|
||||
new_weight = weight.to(devices.device) + updown.to(devices.device)
|
||||
self.weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize)
|
||||
except Exception:
|
||||
# shared.log.error(f'Load network: type=LoRA quant=bnb type={self.quant_type} state={self.quant_state} blocksize={self.blocksize} {e}')
|
||||
extra_network_lora.errors['bnb'] = extra_network_lora.errors.get('bnb', 0) + 1
|
||||
new_weight = None
|
||||
else:
|
||||
try:
|
||||
new_weight = self.weight.to(devices.device) + updown.to(devices.device)
|
||||
except Exception:
|
||||
new_weight = self.weight + updown
|
||||
self.weight = torch.nn.Parameter(new_weight, requires_grad=False)
|
||||
del new_weight
|
||||
if hasattr(self, "qweight") and hasattr(self, "freeze"):
|
||||
|
||||
@@ -288,6 +288,7 @@ def load_scripts():
|
||||
current_basedir = paths.script_path
|
||||
t.record(os.path.basename(scriptfile.basedir) if scriptfile.basedir != paths.script_path else scriptfile.filename)
|
||||
sys.path = syspath
|
||||
|
||||
global scripts_txt2img, scripts_img2img, scripts_control, scripts_postproc # pylint: disable=global-statement
|
||||
scripts_txt2img = ScriptRunner()
|
||||
scripts_img2img = ScriptRunner()
|
||||
|
||||
Reference in New Issue
Block a user