mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
add lora_load_gpu setting
This commit is contained in:
+3
-1
@@ -54,7 +54,8 @@
|
||||
- loras are no longer filtered per detected type vs loaded model type as its unreliable
|
||||
- loras display in networks now shows possible version in top-left corner
|
||||
- correct using of `extra_networks_default_multiplier` if not scale is specified
|
||||
- always keep lora on gpu
|
||||
- setting `lora_load_gpu` to load LoRA directly to GPU
|
||||
*default*: true unless lovwram
|
||||
- **huggingface**:
|
||||
- force logout/login on token change
|
||||
- unified handling of cache folder: set via `HF_HUB` or `HF_HUB_CACHE` or via settings -> system paths
|
||||
@@ -78,6 +79,7 @@
|
||||
- modularize main process loop
|
||||
- massive log cleanup
|
||||
- full lint pass
|
||||
- improve inference mode handling
|
||||
|
||||
|
||||
## Update for 2024-09-13
|
||||
|
||||
@@ -13,6 +13,7 @@ class ModuleTypeLora(network.ModuleType):
|
||||
|
||||
|
||||
class NetworkModuleLora(network.NetworkModule):
|
||||
|
||||
def __init__(self, net: network.Network, weights: network.NetworkWeights):
|
||||
super().__init__(net, weights)
|
||||
self.up_model = self.create_module(weights.w, "lora_up.weight")
|
||||
@@ -21,6 +22,7 @@ class NetworkModuleLora(network.NetworkModule):
|
||||
self.dim = weights.w["lora_down.weight"].shape[0]
|
||||
|
||||
def create_module(self, weights, key, none_ok=False):
|
||||
from modules.shared import opts
|
||||
weight = weights.get(key)
|
||||
if weight is None and none_ok:
|
||||
return None
|
||||
@@ -47,7 +49,8 @@ class NetworkModuleLora(network.NetworkModule):
|
||||
if weight.shape != module.weight.shape:
|
||||
weight = weight.reshape(module.weight.shape)
|
||||
module.weight.copy_(weight)
|
||||
module = module.to(device=devices.device, dtype=devices.dtype)
|
||||
if opts.lora_load_gpu:
|
||||
module = module.to(device=devices.device, dtype=devices.dtype)
|
||||
module.weight.requires_grad_(False)
|
||||
return module
|
||||
|
||||
|
||||
@@ -224,7 +224,7 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt
|
||||
.extra-network-cards .card .overlay .tag { padding: 2px; margin: 2px; background: rgba(70, 70, 70, 0.60); font-size: var(--text-md); cursor: pointer; display: inline-block; }
|
||||
.extra-network-cards .card .actions>span { padding: 4px; font-size: 34px !important; }
|
||||
.extra-network-cards .card .actions>span:hover { color: var(--highlight-color); }
|
||||
.extra-network-cards .card .version { position: absolute; top: 0; left: 0; padding: 2px; font-weight: bolder; text-shadow: 1px 1px black; text-transform: uppercase; font-size: 0.8rem; background: gray; opacity: 75%; margin: 2px; line-height: 0.9rem; }
|
||||
.extra-network-cards .card .version { position: absolute; top: 0; left: 0; padding: 2px; font-weight: bolder; text-shadow: 1px 1px black; text-transform: uppercase; background: gray; opacity: 75%; margin: 4px; line-height: 0.9rem; }
|
||||
.extra-network-cards .card:hover .actions { display: block; }
|
||||
.extra-network-cards .card:hover .overlay .tags { display: block; }
|
||||
.extra-network-cards .card:has(>img[src*="card-no-preview.png"])::before { content: ''; position: absolute; width: 100%; height: 100%; mix-blend-mode: multiply; background-color: var(--data-color); }
|
||||
|
||||
@@ -868,6 +868,7 @@ options_templates.update(options_section(('extra_networks', "Networks"), {
|
||||
"lora_apply_tags": OptionInfo(0, "LoRA auto-apply tags", gr.Slider, {"minimum": -1, "maximum": 32, "step": 1}),
|
||||
"lora_in_memory_limit": OptionInfo(0, "LoRA memory cache", gr.Slider, {"minimum": 0, "maximum": 24, "step": 1}),
|
||||
"lora_functional": OptionInfo(False, "Use Kohya method for handling multiple LoRA", gr.Checkbox, { "visible": False }),
|
||||
"lora_load_gpu": OptionInfo(True if not cmd_opts.lowvram else False, "Load LoRA directly to GPU"),
|
||||
"hypernetwork_enabled": OptionInfo(False, "Enable Hypernetwork support"),
|
||||
"sd_hypernetwork": OptionInfo("None", "Add hypernetwork to prompt", gr.Dropdown, { "choices": ["None"], "visible": False }),
|
||||
"wildcards_enabled": OptionInfo(True, "Enable file wildcards support"),
|
||||
|
||||
@@ -134,18 +134,18 @@ def insert_vectors(embedding, tokenizers, text_encoders, hiddensizes):
|
||||
Future warning, if another text encoder becomes available with embedding dimensions in [768,1280,4096]
|
||||
this may cause collisions.
|
||||
"""
|
||||
for vector, size in zip(embedding.vec, embedding.vector_sizes):
|
||||
if size not in hiddensizes:
|
||||
continue
|
||||
idx = hiddensizes.index(size)
|
||||
unk_token_id = tokenizers[idx].convert_tokens_to_ids(tokenizers[idx].unk_token)
|
||||
if text_encoders[idx].get_input_embeddings().weight.data.shape[0] != len(tokenizers[idx]):
|
||||
text_encoders[idx].resize_token_embeddings(len(tokenizers[idx]))
|
||||
for token, v in zip(embedding.tokens, vector.unbind()):
|
||||
token_id = tokenizers[idx].convert_tokens_to_ids(token)
|
||||
if token_id > unk_token_id:
|
||||
text_encoders[idx].get_input_embeddings().weight.data[token_id] = v
|
||||
|
||||
with devices.inference_context():
|
||||
for vector, size in zip(embedding.vec, embedding.vector_sizes):
|
||||
if size not in hiddensizes:
|
||||
continue
|
||||
idx = hiddensizes.index(size)
|
||||
unk_token_id = tokenizers[idx].convert_tokens_to_ids(tokenizers[idx].unk_token)
|
||||
if text_encoders[idx].get_input_embeddings().weight.data.shape[0] != len(tokenizers[idx]):
|
||||
text_encoders[idx].resize_token_embeddings(len(tokenizers[idx]))
|
||||
for token, v in zip(embedding.tokens, vector.unbind()):
|
||||
token_id = tokenizers[idx].convert_tokens_to_ids(token)
|
||||
if token_id > unk_token_id:
|
||||
text_encoders[idx].get_input_embeddings().weight.data[token_id] = v
|
||||
|
||||
|
||||
class Embedding:
|
||||
@@ -310,6 +310,7 @@ class EmbeddingDatabase:
|
||||
self.register_embedding(embedding, shared.sd_model)
|
||||
except Exception as e:
|
||||
shared.log.error(f'Load embedding: name="{embedding.name}" file="{embedding.filename}" {e}')
|
||||
errors.display(e, f'Load embedding: name="{embedding.name}" file="{embedding.filename}"')
|
||||
return
|
||||
|
||||
def load_from_file(self, path, filename):
|
||||
|
||||
Reference in New Issue
Block a user