Files
modelscope-scepter/scepter/modules/model/embedder/embedder.py
T
2024-01-19 00:44:01 +08:00

673 lines
24 KiB
Python

# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
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
from transformers import CLIPTextModel, CLIPTokenizer
# 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
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 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.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'):
emb_out = embedder(
*[batch[k] for k in embedder.input_keys])
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:
# print("emb.shape", emb.shape)
# print("emb.input_keys", embedder.input_keys)
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
# if "y" in output:
# print("out.shape", output["y"].shape)
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)