lora-direct with bnb

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2024-12-20 18:26:25 -05:00
parent 58ad18ee58
commit dae181fefb
3 changed files with 30 additions and 14 deletions
+3 -2
View File
@@ -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
View File
@@ -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"):
+1
View File
@@ -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()