mirror of
https://github.com/lllyasviel/stable-diffusion-webui-forge.git
synced 2026-07-21 21:01:24 +08:00
llm_adapter
This commit is contained in:
parent
e753e36c8d
commit
436e8758f5
@ -30,7 +30,6 @@ class Anima(ForgeDiffusionEngine):
|
||||
text_encoder=clip.cond_stage_model.qwen3_06b,
|
||||
qwen_tokenizer=clip.tokenizer.qwen3_06b,
|
||||
t5_tokenizer=clip.tokenizer.t5xxl,
|
||||
unet=unet,
|
||||
)
|
||||
|
||||
self.forge_objects = ForgeObjects(unet=unet, clip=clip, vae=vae, clipvision=None)
|
||||
|
||||
@ -631,6 +631,15 @@ def preprocess_state_dict(sd: dict[str, torch.Tensor]) -> dict[str, torch.Tensor
|
||||
return sd
|
||||
|
||||
|
||||
def process_anima(dit: dict[str, torch.Tensor], enc: dict[str, torch.Tensor]):
|
||||
# move LLMAdapter from transformer to text_encoder
|
||||
|
||||
keys = list(dit.keys())
|
||||
for k in keys:
|
||||
if k.startswith("llm_adapter"):
|
||||
enc[k] = dit.pop(k)
|
||||
|
||||
|
||||
def split_state_dict(sd, additional_state_dicts: list = None):
|
||||
import huggingface_guess
|
||||
|
||||
@ -672,6 +681,9 @@ def split_state_dict(sd, additional_state_dicts: list = None):
|
||||
|
||||
state_dict["ignore"] = sd
|
||||
|
||||
if "Anima" in guess.huggingface_repo:
|
||||
process_anima(state_dict["transformer"], state_dict["text_encoder"])
|
||||
|
||||
print_dict = {k: len(v) for k, v in state_dict.items()}
|
||||
logger.debug(f"StateDict Keys: {print_dict}")
|
||||
|
||||
|
||||
@ -791,14 +791,4 @@ class LLMAdapter(nn.Module):
|
||||
|
||||
|
||||
class Anima(MiniTrainDIT):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.llm_adapter = LLMAdapter()
|
||||
|
||||
@torch.inference_mode()
|
||||
def preprocess_text_embeds(self, text_embeds, text_ids):
|
||||
if text_ids is not None:
|
||||
device = self.llm_adapter.embed.weight.device
|
||||
return self.llm_adapter(text_embeds.to(device), text_ids.to(device))
|
||||
else:
|
||||
return text_embeds
|
||||
pass
|
||||
|
||||
@ -14,6 +14,7 @@ if pytorch_attention_enabled:
|
||||
else:
|
||||
from backend.attention import attention_basic as attention_function
|
||||
|
||||
from backend.nn.anima import LLMAdapter
|
||||
from backend.nn.llm import qwen_vl
|
||||
|
||||
|
||||
@ -513,6 +514,14 @@ class Qwen3_06B(BaseLlama, nn.Module):
|
||||
|
||||
self.model = Llama2_(config)
|
||||
|
||||
self.llm_adapter = LLMAdapter()
|
||||
|
||||
def preprocess_text_embeds(self, text_embeds, text_ids):
|
||||
if text_ids is not None:
|
||||
return self.llm_adapter(text_embeds, text_ids)
|
||||
else:
|
||||
return text_embeds
|
||||
|
||||
|
||||
class Qwen3_4B(BaseLlama, nn.Module):
|
||||
def __init__(self, config_dict):
|
||||
|
||||
@ -1,9 +1,7 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from backend.patcher.unet import UnetPatcher
|
||||
|
||||
import weakref
|
||||
from backend.nn.llm.llama import Qwen3_06B
|
||||
|
||||
import torch
|
||||
|
||||
@ -21,22 +19,16 @@ class PromptChunk:
|
||||
|
||||
|
||||
class AnimaTextProcessingEngine:
|
||||
def __init__(self, text_encoder, qwen_tokenizer, t5_tokenizer, unet):
|
||||
def __init__(self, text_encoder, qwen_tokenizer, t5_tokenizer):
|
||||
super().__init__()
|
||||
|
||||
self.text_encoder = text_encoder
|
||||
self.text_encoder: "Qwen3_06B" = text_encoder
|
||||
self.qwen_tokenizer = qwen_tokenizer
|
||||
self.t5_tokenizer = t5_tokenizer
|
||||
|
||||
self._unet = weakref.ref(unet)
|
||||
|
||||
self.id_pad = 151643
|
||||
self.id_end = 1
|
||||
|
||||
@property
|
||||
def unet(self) -> " UnetPatcher":
|
||||
return self._unet()
|
||||
|
||||
def tokenize(self, texts):
|
||||
return (
|
||||
self.qwen_tokenizer(texts, truncation=False, add_special_tokens=False)["input_ids"],
|
||||
@ -116,12 +108,12 @@ class AnimaTextProcessingEngine:
|
||||
return zs
|
||||
|
||||
def anima_preprocess(self, cross_attn: torch.Tensor, t5xxl_ids: torch.Tensor, t5xxl_weights: torch.Tensor) -> torch.Tensor:
|
||||
dtype: torch.dtype = self.unet.model.computation_dtype
|
||||
device = memory_management.text_encoder_device()
|
||||
|
||||
cross_attn = cross_attn.unsqueeze(0).to(dtype=dtype)
|
||||
t5xxl_ids = t5xxl_ids.unsqueeze(0)
|
||||
cross_attn = cross_attn.unsqueeze(0).to(device=device)
|
||||
t5xxl_ids = t5xxl_ids.unsqueeze(0).to(device=device)
|
||||
|
||||
cross_attn = self.unet.model.diffusion_model.preprocess_text_embeds(cross_attn, t5xxl_ids)
|
||||
cross_attn = self.text_encoder.preprocess_text_embeds(cross_attn, t5xxl_ids)
|
||||
if t5xxl_weights is not None:
|
||||
cross_attn *= t5xxl_weights.unsqueeze(0).unsqueeze(-1).to(cross_attn)
|
||||
|
||||
@ -161,23 +153,6 @@ class AnimaTextProcessingEngine:
|
||||
index = 0
|
||||
embeds_info = []
|
||||
|
||||
for o in other_embeds:
|
||||
emb, extra = self.text_encoder.preprocess_embed(o[1], device=device)
|
||||
if emb is None:
|
||||
index += -1
|
||||
continue
|
||||
|
||||
ind = index + o[0]
|
||||
emb = emb.view(1, -1, emb.shape[-1]).to(device=device, dtype=torch.float32)
|
||||
emb_shape = emb.shape[1]
|
||||
|
||||
assert emb.shape[-1] == tokens_embed.shape[-1]
|
||||
tokens_embed = torch.cat([tokens_embed[:, :ind], emb, tokens_embed[:, ind:]], dim=1)
|
||||
attention_mask = attention_mask[:ind] + [1] * emb_shape + attention_mask[ind:]
|
||||
index += emb_shape - 1
|
||||
emb_type = o[1].get("type", None)
|
||||
embeds_info.append({"type": emb_type, "index": ind, "size": emb_shape, "extra": extra})
|
||||
|
||||
embeds_out.append(tokens_embed)
|
||||
attention_masks.append(attention_mask)
|
||||
num_tokens.append(sum(attention_mask))
|
||||
|
||||
Loading…
Reference in New Issue
Block a user