783 lines
28 KiB
Python
783 lines
28 KiB
Python
# -*- coding: utf-8 -*-
|
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
|
import warnings
|
|
from collections import OrderedDict
|
|
from contextlib import nullcontext
|
|
from typing import Dict
|
|
|
|
import numpy as np
|
|
import open_clip
|
|
import torch
|
|
import torch.nn as nn
|
|
import torch.utils.dlpack
|
|
from einops import rearrange
|
|
from torch.utils.checkpoint import checkpoint
|
|
|
|
# to check
|
|
from scepter.modules.model.backbone.unet.unet_utils import Timestep
|
|
from scepter.modules.model.registry import EMBEDDERS
|
|
from scepter.modules.model.utils.basic_utils import expand_dims_like
|
|
from scepter.modules.utils.config import dict_to_yaml
|
|
from scepter.modules.utils.distribute import we
|
|
from scepter.modules.utils.file_system import FS
|
|
|
|
from .base_embedder import BaseEmbedder
|
|
from .resampler import Resampler
|
|
|
|
try:
|
|
from transformers import CLIPTextModel, CLIPTokenizer, CLIPVisionModelWithProjection
|
|
except Exception as e:
|
|
warnings.warn(
|
|
f'Import transformers error, please deal with this problem: {e}')
|
|
|
|
|
|
def autocast(f, enabled=True):
|
|
def do_autocast(*args, **kwargs):
|
|
with torch.cuda.amp.autocast(
|
|
enabled=enabled,
|
|
dtype=torch.get_autocast_gpu_dtype(),
|
|
cache_enabled=torch.is_autocast_cache_enabled(),
|
|
):
|
|
return f(*args, **kwargs)
|
|
|
|
return do_autocast
|
|
|
|
|
|
@EMBEDDERS.register_class()
|
|
class FrozenCLIPEmbedder(BaseEmbedder):
|
|
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
|
para_dict = {
|
|
'PRETRAINED_MODEL': {
|
|
'value': None,
|
|
'description': 'You should set pretrained_model: modelcard.'
|
|
},
|
|
'TOKENIZER_PATH': {
|
|
'value':
|
|
None,
|
|
'description':
|
|
'If you want to use tokenizer in embedder, you should set this field.'
|
|
},
|
|
'MAX_LENGTH': {
|
|
'value': 77,
|
|
'description': ''
|
|
},
|
|
'FREEZE': {
|
|
'value': True,
|
|
'description': ''
|
|
},
|
|
'USE_GRAD': {
|
|
'value': False,
|
|
'description': 'Compute grad or not.'
|
|
},
|
|
'LAYER': {
|
|
'value': 'last',
|
|
'description': ''
|
|
},
|
|
'LAYER_IDX': {
|
|
'value': None,
|
|
'description': ''
|
|
},
|
|
'USE_FINAL_LAYER_NORM': {
|
|
'value': False,
|
|
'description': ''
|
|
},
|
|
}
|
|
LAYERS = ['last', 'pooled', 'hidden']
|
|
|
|
def __init__(self, cfg, logger=None):
|
|
super().__init__(cfg, logger=logger)
|
|
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
|
if tokenizer_path is not None:
|
|
with FS.get_dir_to_local_dir(tokenizer_path,
|
|
wait_finish=True) as local_path:
|
|
self.tokenizer = CLIPTokenizer.from_pretrained(local_path)
|
|
|
|
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
|
if pretrained_model is None:
|
|
raise 'You should set pretrained_model: modelcard.'
|
|
with FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL,
|
|
wait_finish=True) as local_path:
|
|
self.transformer = CLIPTextModel.from_pretrained(local_path)
|
|
|
|
self.use_grad = cfg.get('USE_GRAD', False)
|
|
self.freeze_flag = cfg.get('FREEZE', True)
|
|
if self.freeze_flag:
|
|
self.freeze()
|
|
|
|
self.max_length = cfg.get('MAX_LENGTH', 77)
|
|
self.layer = cfg.get('LAYER', 'last')
|
|
self.layer_idx = cfg.get('LAYER_IDX', None)
|
|
self.use_final_layer_norm = cfg.get('USE_FINAL_LAYER_NORM', False)
|
|
assert self.layer in self.LAYERS
|
|
if self.layer == 'hidden':
|
|
assert self.layer_idx is not None
|
|
assert 0 <= abs(self.layer_idx) <= 12
|
|
|
|
def freeze(self):
|
|
self.transformer = self.transformer.eval()
|
|
for param in self.parameters():
|
|
param.requires_grad = False
|
|
|
|
# @torch.no_grad()
|
|
def _forward(self, text):
|
|
batch_encoding = self.tokenizer(text,
|
|
truncation=True,
|
|
max_length=self.max_length,
|
|
return_length=True,
|
|
return_overflowing_tokens=False,
|
|
padding='max_length',
|
|
return_tensors='pt')
|
|
tokens = batch_encoding['input_ids'].to(we.device_id)
|
|
outputs = self.transformer(input_ids=tokens,
|
|
output_hidden_states=self.layer == 'hidden')
|
|
if self.layer == 'last':
|
|
z = outputs.last_hidden_state
|
|
elif self.layer == 'pooled':
|
|
z = outputs.pooler_output[:, None, :]
|
|
if self.use_final_layer_norm:
|
|
z = self.transformer.text_model.final_layer_norm(z)
|
|
else:
|
|
z = outputs.hidden_states[self.layer_idx]
|
|
if self.use_final_layer_norm:
|
|
z = self.transformer.text_model.final_layer_norm(z)
|
|
return z
|
|
|
|
@autocast
|
|
def forward(self, text):
|
|
if not self.use_grad:
|
|
with torch.no_grad():
|
|
output = self._forward(text)
|
|
else:
|
|
output = self._forward(text)
|
|
return output
|
|
|
|
def encode(self, text):
|
|
return self(text)
|
|
|
|
# @torch.no_grad()
|
|
def _encode_text(self,
|
|
tokens,
|
|
tokenizer=None,
|
|
append_sentence_embedding=False):
|
|
outputs = self.transformer(input_ids=tokens,
|
|
output_hidden_states=self.layer == 'hidden')
|
|
if self.layer == 'last':
|
|
z = outputs.last_hidden_state
|
|
elif self.layer == 'pooled':
|
|
z = outputs.pooler_output[:, None, :]
|
|
if self.use_final_layer_norm:
|
|
z = self.transformer.text_model.final_layer_norm(z)
|
|
else:
|
|
z = outputs.hidden_states[self.layer_idx]
|
|
if self.use_final_layer_norm:
|
|
z = self.transformer.text_model.final_layer_norm(z)
|
|
return z
|
|
|
|
def encode_text(self,
|
|
tokens,
|
|
tokenizer=None,
|
|
append_sentence_embedding=False):
|
|
if not self.use_grad:
|
|
with torch.no_grad():
|
|
output = self._encode_text(tokens, tokenizer,
|
|
append_sentence_embedding)
|
|
else:
|
|
output = self._encode_text(tokens, tokenizer,
|
|
append_sentence_embedding)
|
|
return output
|
|
|
|
@staticmethod
|
|
def get_config_template():
|
|
return dict_to_yaml('MODELS',
|
|
__class__.__name__,
|
|
FrozenCLIPEmbedder.para_dict,
|
|
set_name=True)
|
|
|
|
|
|
@EMBEDDERS.register_class()
|
|
class FrozenOpenCLIPEmbedder(BaseEmbedder):
|
|
"""
|
|
Uses the OpenCLIP transformer encoder for text
|
|
"""
|
|
para_dict = {
|
|
'ARCH': {
|
|
'value': 'ViT-H-14',
|
|
'description': ''
|
|
},
|
|
'PRETRAINED_MODEL': {
|
|
'value': '',
|
|
'description': ''
|
|
},
|
|
'MAX_LENGTH': {
|
|
'value': 77,
|
|
'description': ''
|
|
},
|
|
'FREEZE': {
|
|
'value': True,
|
|
'description': ''
|
|
},
|
|
'USE_GRAD': {
|
|
'value': False,
|
|
'description': 'Compute grad or not.'
|
|
},
|
|
'LAYER': {
|
|
'value': 'last',
|
|
'description': ''
|
|
},
|
|
}
|
|
LAYERS = ['last', 'penultimate']
|
|
|
|
def __init__(self, cfg, logger=None):
|
|
super().__init__(cfg, logger=logger)
|
|
|
|
arch = cfg.get('ARCH', 'ViT-H-14')
|
|
model, _, _ = open_clip.create_model_and_transforms(
|
|
arch, device=torch.device('cpu'), pretrained=None)
|
|
del model.visual
|
|
if cfg.PRETRAINED_MODEL is not None:
|
|
with FS.get_from(cfg.PRETRAINED_MODEL,
|
|
wait_finish=True) as local_path:
|
|
model.load_state_dict(torch.load(local_path), strict=False)
|
|
self.model = model
|
|
|
|
self.use_grad = cfg.get('USE_GRAD', False)
|
|
self.freeze_flag = cfg.get('FREEZE', True)
|
|
if self.freeze_flag:
|
|
self.freeze()
|
|
|
|
self.max_length = cfg.get('MAX_LENGTH', 77)
|
|
self.layer = cfg.get('LAYER', 'penultimate')
|
|
assert self.layer in self.LAYERS
|
|
if self.layer == 'last':
|
|
self.layer_idx = 0
|
|
elif self.layer == 'penultimate':
|
|
self.layer_idx = 1
|
|
else:
|
|
raise NotImplementedError()
|
|
|
|
def freeze(self):
|
|
self.model = self.model.eval()
|
|
for param in self.parameters():
|
|
param.requires_grad = False
|
|
|
|
@autocast
|
|
def forward(self, text):
|
|
embedding_context = nullcontext if self.use_grad else torch.no_grad
|
|
with embedding_context():
|
|
tokens = open_clip.tokenize(text)
|
|
z = self.encode_with_transformer(tokens.to(we.device_id))
|
|
return z
|
|
|
|
def encode_with_transformer(self, text):
|
|
embedding_context = nullcontext if self.use_grad else torch.no_grad
|
|
with embedding_context():
|
|
x = self.model.token_embedding(
|
|
text) # [batch_size, n_ctx, d_model]
|
|
x = x + self.model.positional_embedding
|
|
x = x.permute(1, 0, 2) # NLD -> LND
|
|
x = self.text_transformer_forward(x,
|
|
attn_mask=self.model.attn_mask)
|
|
x = x.permute(1, 0, 2) # LND -> NLD
|
|
x = self.model.ln_final(x)
|
|
return x
|
|
|
|
def text_transformer_forward(self, x: torch.Tensor, attn_mask=None):
|
|
embedding_context = nullcontext if self.use_grad else torch.no_grad
|
|
with embedding_context():
|
|
for i, r in enumerate(self.model.transformer.resblocks):
|
|
if i == len(self.model.transformer.resblocks) - self.layer_idx:
|
|
break
|
|
if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting(
|
|
):
|
|
x = checkpoint(r, x, attn_mask)
|
|
else:
|
|
x = r(x, attn_mask=attn_mask)
|
|
return x
|
|
|
|
def encode_text(self,
|
|
tokens,
|
|
tokenizer=None,
|
|
append_sentence_embedding=False):
|
|
embedding_context = nullcontext if self.use_grad else torch.no_grad
|
|
with embedding_context():
|
|
z = self.encode_with_transformer(tokens.to(we.device_id))
|
|
return z
|
|
|
|
def encode(self, text):
|
|
return self(text)
|
|
|
|
@staticmethod
|
|
def get_config_template():
|
|
return dict_to_yaml('MODELS',
|
|
__class__.__name__,
|
|
FrozenOpenCLIPEmbedder.para_dict,
|
|
set_name=True)
|
|
|
|
|
|
@EMBEDDERS.register_class()
|
|
class FrozenOpenCLIPEmbedder2(BaseEmbedder):
|
|
"""
|
|
Uses the OpenCLIP transformer encoder for text
|
|
"""
|
|
"""
|
|
Uses the OpenCLIP transformer encoder for text
|
|
"""
|
|
para_dict = {
|
|
'ARCH': {
|
|
'value': 'ViT-H-14',
|
|
'description': ''
|
|
},
|
|
'PRETRAINED_MODEL': {
|
|
'value': 'laion2b_s32b_b79k',
|
|
'description': ''
|
|
},
|
|
'MAX_LENGTH': {
|
|
'value': 77,
|
|
'description': ''
|
|
},
|
|
'FREEZE': {
|
|
'value': True,
|
|
'description': ''
|
|
},
|
|
'USE_GRAD': {
|
|
'value': False,
|
|
'description': 'Compute grad or not.'
|
|
},
|
|
'ALWAYS_RETURN_POOLED': {
|
|
'value':
|
|
False,
|
|
'description':
|
|
'Whether always return pooled results or not ,default False.'
|
|
},
|
|
'LEGACY': {
|
|
'value':
|
|
True,
|
|
'description':
|
|
'Whether use legacy returnd feature or not ,default True.'
|
|
},
|
|
'LAYER': {
|
|
'value': 'last',
|
|
'description': ''
|
|
},
|
|
}
|
|
|
|
LAYERS = ['pooled', 'last', 'penultimate']
|
|
|
|
def __init__(self, cfg, logger=None):
|
|
super().__init__(cfg, logger=logger)
|
|
arch = cfg.get('ARCH', 'ViT-H-14')
|
|
if cfg.get('PRETRAINED_MODEL', None) is None:
|
|
model, _, _ = open_clip.create_model_and_transforms(
|
|
arch, device=torch.device('cpu'), pretrained=None)
|
|
del model.visual
|
|
else:
|
|
with FS.get_from(cfg.PRETRAINED_MODEL,
|
|
wait_finish=True) as local_path:
|
|
model, _, _ = open_clip.create_model_and_transforms(
|
|
arch, device=torch.device('cpu'), pretrained=local_path)
|
|
del model.visual
|
|
self.model = model
|
|
|
|
self.max_length = cfg.get('MAX_LENGTH', 77)
|
|
self.layer = cfg.get('LAYER', 'last')
|
|
self.return_pooled = cfg.get('ALWAYS_RETURN_POOLED', False)
|
|
self.use_grad = cfg.get('USE_GRAD', False)
|
|
self.freeze_flag = cfg.get('FREEZE', True)
|
|
if self.freeze_flag:
|
|
self.freeze()
|
|
if self.layer == 'last':
|
|
self.layer_idx = 0
|
|
elif self.layer == 'penultimate':
|
|
self.layer_idx = 1
|
|
else:
|
|
raise NotImplementedError()
|
|
self.legacy = cfg.get('LEGACY', True)
|
|
|
|
def freeze(self):
|
|
self.model = self.model.eval()
|
|
for param in self.parameters():
|
|
param.requires_grad = False
|
|
|
|
@autocast
|
|
def forward(self, text):
|
|
embedding_context = nullcontext if self.use_grad else torch.no_grad
|
|
with embedding_context():
|
|
tokens = open_clip.tokenize(text)
|
|
z = self.encode_with_transformer(tokens.to(we.device_id))
|
|
if not self.return_pooled and self.legacy:
|
|
return z
|
|
if self.return_pooled:
|
|
assert not self.legacy
|
|
return z[self.layer], z['pooled']
|
|
return z[self.layer]
|
|
|
|
def encode_with_transformer(self, text):
|
|
embedding_context = nullcontext if self.use_grad else torch.no_grad
|
|
with embedding_context():
|
|
x = self.model.token_embedding(
|
|
text) # [batch_size, n_ctx, d_model]
|
|
x = x + self.model.positional_embedding
|
|
x = x.permute(1, 0, 2) # NLD -> LND
|
|
x = self.text_transformer_forward(x,
|
|
attn_mask=self.model.attn_mask)
|
|
if self.legacy:
|
|
x = x[self.layer]
|
|
x = self.model.ln_final(x)
|
|
return x
|
|
else:
|
|
# x is a dict and will stay a dict
|
|
o = x['last']
|
|
o = self.model.ln_final(o)
|
|
pooled = self.pool(o, text)
|
|
x['pooled'] = pooled
|
|
return x
|
|
|
|
def pool(self, x, text):
|
|
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
|
x = (x[torch.arange(x.shape[0]),
|
|
text.argmax(dim=-1)] @ self.model.text_projection)
|
|
return x
|
|
|
|
def text_transformer_forward(self, x: torch.Tensor, attn_mask=None):
|
|
outputs = {}
|
|
embedding_context = nullcontext if self.use_grad else torch.no_grad
|
|
with embedding_context():
|
|
for i, r in enumerate(self.model.transformer.resblocks):
|
|
if i == len(self.model.transformer.resblocks) - 1:
|
|
outputs['penultimate'] = x.permute(1, 0, 2) # LND -> NLD
|
|
if (self.model.transformer.grad_checkpointing
|
|
and not torch.jit.is_scripting()):
|
|
x = checkpoint(r, x, attn_mask)
|
|
else:
|
|
x = r(x, attn_mask=attn_mask)
|
|
outputs['last'] = x.permute(1, 0, 2) # LND -> NLD
|
|
return outputs
|
|
|
|
def encode(self, text):
|
|
return self(text)
|
|
|
|
def encode_text(self,
|
|
tokens,
|
|
tokenizer=None,
|
|
append_sentence_embedding=False):
|
|
z = self.encode_with_transformer(tokens.to(we.device_id))
|
|
if not self.return_pooled and self.legacy:
|
|
return z
|
|
if self.return_pooled:
|
|
assert not self.legacy
|
|
return z[self.layer], z['pooled']
|
|
return z[self.layer]
|
|
|
|
@staticmethod
|
|
def get_config_template():
|
|
return dict_to_yaml('MODELS',
|
|
__class__.__name__,
|
|
FrozenOpenCLIPEmbedder2.para_dict,
|
|
set_name=True)
|
|
|
|
|
|
@EMBEDDERS.register_class()
|
|
class ConcatTimestepEmbedderND(BaseEmbedder):
|
|
"""embeds each dimension independently and concatenates them"""
|
|
para_dict = {
|
|
'OUT_DIM': {
|
|
'value': 256,
|
|
'description': 'Output dim'
|
|
},
|
|
}
|
|
|
|
def __init__(self, cfg, logger=None):
|
|
super().__init__(cfg=cfg, logger=logger)
|
|
outdim = cfg.get('OUT_DIM', 256)
|
|
self.timestep = Timestep(outdim, legacy=True)
|
|
self.outdim = outdim
|
|
|
|
def forward(self, x):
|
|
if x.ndim == 1:
|
|
x = x[:, None]
|
|
assert len(x.shape) == 2
|
|
b, dims = x.shape[0], x.shape[1]
|
|
x = rearrange(x, 'b d -> (b d)')
|
|
emb = self.timestep(x)
|
|
emb = rearrange(emb,
|
|
'(b d) d2 -> b (d d2)',
|
|
b=b,
|
|
d=dims,
|
|
d2=self.outdim)
|
|
return emb
|
|
|
|
@staticmethod
|
|
def get_config_template():
|
|
return dict_to_yaml('MODELS',
|
|
__class__.__name__,
|
|
ConcatTimestepEmbedderND.para_dict,
|
|
set_name=True)
|
|
|
|
|
|
@EMBEDDERS.register_class()
|
|
class IPAdapterPlusEmbedder(BaseEmbedder):
|
|
def __init__(self, cfg, logger=None):
|
|
super().__init__(cfg, logger=logger)
|
|
|
|
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
|
|
wait_finish=True) as local_path:
|
|
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
|
local_path)
|
|
|
|
self.image_proj_model = Resampler(
|
|
dim=self.cfg.get('IN_DIM', 768),
|
|
depth=self.cfg.get('DEPTH', 4),
|
|
dim_head=64,
|
|
heads=self.cfg.get('HEADS', 12),
|
|
num_queries=self.cfg.get('NUM_TOKENS', 16),
|
|
embedding_dim=self.image_encoder.config.hidden_size,
|
|
output_dim=self.cfg.get('CROSSATTN_DIM', 768),
|
|
ff_mult=4,
|
|
)
|
|
|
|
with FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) as local_path:
|
|
ckpt = torch.load(local_path, map_location='cpu')
|
|
self.image_proj_model.load_state_dict(ckpt['image_proj'],
|
|
strict=True)
|
|
|
|
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
|
|
self.cfg.get('CROSSATTN_DIM', 768))
|
|
|
|
def encode(self, ref_ip, ref_detail):
|
|
encoder_output = self.image_encoder(ref_ip, output_hidden_states=True)
|
|
image_prompt_embeds = self.image_proj_model(
|
|
encoder_output.hidden_states[-2])
|
|
encoder_output_2 = self.image_encoder(ref_detail,
|
|
output_hidden_states=True)
|
|
image_patch_embeds = self.patch_projector(
|
|
encoder_output_2.last_hidden_state)
|
|
out = {
|
|
'img_crossattn': image_prompt_embeds,
|
|
'ref_crossattn': image_patch_embeds,
|
|
}
|
|
return out
|
|
|
|
def forward(self, ref_ip, ref_detail):
|
|
return self.encode(ref_ip, ref_detail)
|
|
|
|
|
|
class RefCrossEmbedder(BaseEmbedder):
|
|
def __init__(self, cfg, logger=None):
|
|
super().__init__(cfg, logger=logger)
|
|
|
|
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
|
|
wait_finish=True) as local_path:
|
|
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
|
local_path)
|
|
|
|
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
|
|
self.cfg.get('CROSSATTN_DIM', 768))
|
|
|
|
def encode(self, img):
|
|
encoder_output = self.image_encoder(img, output_hidden_states=True)
|
|
image_patch_embeds = self.patch_projector(
|
|
encoder_output.last_hidden_state)
|
|
out = {
|
|
'ref_crossattn': image_patch_embeds,
|
|
}
|
|
return out
|
|
|
|
def forward(self, img):
|
|
return self.encode(img)
|
|
|
|
|
|
@EMBEDDERS.register_class()
|
|
class TransparentEmbedder(BaseEmbedder):
|
|
def forward(self, *args):
|
|
out = dict()
|
|
for key, val in zip(self.input_keys, args):
|
|
out[key] = val
|
|
return out
|
|
|
|
|
|
@EMBEDDERS.register_class()
|
|
class NoiseConcatEmbedder(BaseEmbedder):
|
|
def forward(self, *args):
|
|
return {'concat': torch.cat(args, dim=1)}
|
|
|
|
|
|
@EMBEDDERS.register_class()
|
|
class GeneralConditioner(BaseEmbedder):
|
|
OUTPUT_DIM2KEYS = {2: 'y', 3: 'crossattn', 4: 'concat', 5: 'concat'}
|
|
KEY2CATDIM = {'y': 1, 'crossattn': 2, 'concat': 1}
|
|
para_dict = {
|
|
'EMBEDDERS': [],
|
|
'USE_GRAD': {
|
|
'value': False,
|
|
'description': 'Compute grad or not.'
|
|
},
|
|
}
|
|
para_dict.update(para_dict)
|
|
|
|
def __init__(self, cfg, logger=None):
|
|
super().__init__(cfg, logger=logger)
|
|
emb_models = cfg.get('EMBEDDERS', [])
|
|
use_grad = cfg.get('USE_GRAD', False)
|
|
self.embedders = nn.ModuleList([])
|
|
for n, embconfig in enumerate(emb_models):
|
|
embconfig.USE_GRAD = use_grad if not embconfig.have(
|
|
'USE_GRAD'
|
|
) or embconfig.USE_GRAD is None else embconfig.USE_GRAD
|
|
embedder = EMBEDDERS.build(embconfig, logger=logger)
|
|
embedder.ucg_rate = embconfig.get('UCG_RATE', 0.0)
|
|
embedder.input_keys = embconfig.get('INPUT_KEYS', [])
|
|
embedder.legacy_ucg_val = embconfig.get('LEGACY_UCG_VALUE', None)
|
|
if embedder.legacy_ucg_val is not None:
|
|
embedder.ucg_prng = np.random.RandomState()
|
|
|
|
self.embedders.append(embedder)
|
|
|
|
def load_pretrained_model(self, pretrained_model):
|
|
if pretrained_model is not None:
|
|
with FS.get_from(pretrained_model,
|
|
wait_finish=True) as local_model:
|
|
self.init_from_ckpt(local_model)
|
|
|
|
def init_from_ckpt(self, path, ignore_keys=list()):
|
|
if path.endswith('safetensors'):
|
|
from safetensors.torch import load_file as load_safetensors
|
|
sd = load_safetensors(path)
|
|
else:
|
|
sd = torch.load(path, map_location='cpu')
|
|
new_sd = OrderedDict()
|
|
for k, v in sd.items():
|
|
ignored = False
|
|
for ik in ignore_keys:
|
|
if ik in k:
|
|
if we.rank == 0:
|
|
self.logger.info(
|
|
'Ignore key {} from state_dict.'.format(k))
|
|
ignored = True
|
|
break
|
|
if not ignored:
|
|
new_sd[k] = v
|
|
|
|
missing, unexpected = self.load_state_dict(new_sd, strict=False)
|
|
if we.rank == 0:
|
|
self.logger.info(
|
|
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
|
)
|
|
if len(missing) > 0:
|
|
self.logger.info(f'Missing Keys:\n {missing}')
|
|
if len(unexpected) > 0:
|
|
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
|
|
|
def possibly_get_ucg_val(self, embedder, batch: Dict) -> Dict:
|
|
assert embedder.legacy_ucg_val is not None
|
|
p = embedder.ucg_rate
|
|
val = embedder.legacy_ucg_val
|
|
for i in range(len(batch[embedder.input_key])):
|
|
if embedder.ucg_prng.choice(2, p=[1 - p, p]):
|
|
batch[embedder.input_key][i] = val
|
|
return batch
|
|
|
|
def forward(self, batch: Dict, force_zero_embeddings=None) -> Dict:
|
|
output = dict()
|
|
if force_zero_embeddings is None:
|
|
force_zero_embeddings = []
|
|
for embedder in self.embedders:
|
|
embedding_context = nullcontext if hasattr(
|
|
embedder, 'use_grad') and embedder.use_grad else torch.no_grad
|
|
with embedding_context():
|
|
if hasattr(embedder, 'input_key') and (embedder.input_key
|
|
is not None):
|
|
if embedder.input_key not in batch:
|
|
continue
|
|
if embedder.legacy_ucg_val is not None:
|
|
batch = self.possibly_get_ucg_val(embedder, batch)
|
|
emb_out = embedder(batch[embedder.input_key])
|
|
elif hasattr(embedder, 'input_keys'):
|
|
if any([k not in batch for k in embedder.input_keys]):
|
|
continue
|
|
emb_out = embedder(
|
|
*[batch[k] for k in embedder.input_keys])
|
|
|
|
if isinstance(emb_out, dict):
|
|
for key, val in emb_out.items():
|
|
if key in output:
|
|
assert key in self.KEY2CATDIM
|
|
output[key] = torch.cat([output[key], val],
|
|
dim=self.KEY2CATDIM[key])
|
|
else:
|
|
output[key] = val
|
|
else:
|
|
assert isinstance(
|
|
emb_out, (torch.Tensor, list, tuple)
|
|
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
|
|
|
|
if not isinstance(emb_out, (list, tuple)):
|
|
emb_out = [emb_out]
|
|
|
|
for emb in emb_out:
|
|
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
|
|
|
|
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
|
|
emb = (expand_dims_like(
|
|
torch.bernoulli(
|
|
(1.0 - embedder.ucg_rate) *
|
|
torch.ones(emb.shape[0], device=emb.device)),
|
|
emb,
|
|
) * emb)
|
|
|
|
if (hasattr(embedder, 'input_keys')):
|
|
if np.sum(
|
|
np.array([
|
|
key in force_zero_embeddings
|
|
for key in embedder.input_keys
|
|
])) > 0:
|
|
emb = torch.zeros_like(emb)
|
|
|
|
if out_key in output:
|
|
output[out_key] = torch.cat((output[out_key], emb),
|
|
self.KEY2CATDIM[out_key])
|
|
else:
|
|
output[out_key] = emb
|
|
|
|
return output
|
|
|
|
def get_unconditional_conditioning(self,
|
|
batch_c,
|
|
batch_uc=None,
|
|
force_uc_zero_embeddings=None):
|
|
if force_uc_zero_embeddings is None:
|
|
force_uc_zero_embeddings = []
|
|
ucg_rates = list()
|
|
for embedder in self.embedders:
|
|
ucg_rates.append(embedder.ucg_rate)
|
|
embedder.ucg_rate = 0.0
|
|
c = self(batch_c)
|
|
uc = self(batch_c if batch_uc is None else batch_uc,
|
|
force_uc_zero_embeddings)
|
|
|
|
for embedder, rate in zip(self.embedders, ucg_rates):
|
|
embedder.ucg_rate = rate
|
|
return c, uc
|
|
|
|
def encode(self,
|
|
batch_dict,
|
|
is_unconditional=False,
|
|
force_uc_zero_embeddings=None):
|
|
ucg_rates = list()
|
|
for embedder in self.embedders:
|
|
ucg_rates.append(embedder.ucg_rate)
|
|
embedder.ucg_rate = 0.0
|
|
if is_unconditional:
|
|
c = self(batch_dict, force_uc_zero_embeddings)
|
|
else:
|
|
c = self(batch_dict)
|
|
for embedder, rate in zip(self.embedders, ucg_rates):
|
|
embedder.ucg_rate = rate
|
|
return c
|
|
|
|
@staticmethod
|
|
def get_config_template():
|
|
return dict_to_yaml('MODELS',
|
|
__class__.__name__,
|
|
GeneralConditioner.para_dict,
|
|
set_name=True)
|