Binary file not shown.
|
Before Width: | Height: | Size: 118 KiB After Width: | Height: | Size: 109 KiB |
@@ -5,18 +5,14 @@ import os.path
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL.Image import Image
|
||||
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.model.network.diffusion.schedules import noise_schedule
|
||||
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
|
||||
TOKENIZERS)
|
||||
from scepter.modules.model.utils.data_utils import crop_back
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.utils.env import get_available_memory
|
||||
|
||||
from .diffusion_inference import DiffusionInference
|
||||
|
||||
|
||||
def get_model(model_tuple):
|
||||
@@ -24,7 +20,7 @@ def get_model(model_tuple):
|
||||
return model_tuple['model']
|
||||
|
||||
|
||||
class LargenInference():
|
||||
class LargenInference(DiffusionInference):
|
||||
'''
|
||||
define vae, unet, text-encoder, tuner, refiner components
|
||||
support to load the components dynamicly.
|
||||
@@ -37,45 +33,6 @@ class LargenInference():
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
|
||||
def init_from_cfg(self, cfg):
|
||||
self.name = cfg.NAME
|
||||
self.is_default = cfg.get('IS_DEFAULT', False)
|
||||
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
|
||||
assert cfg.have('MODEL')
|
||||
cfg.MODEL = self.redefine_paras(cfg.MODEL)
|
||||
self.diffusion = self.load_schedule(cfg.MODEL.SCHEDULE)
|
||||
self.diffusion_model = self.infer_model(
|
||||
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
|
||||
'DIFFUSION_MODEL',
|
||||
None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None
|
||||
self.first_stage_model = self.infer_model(
|
||||
cfg.MODEL.FIRST_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'FIRST_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None
|
||||
self.cond_stage_model = self.infer_model(
|
||||
cfg.MODEL.COND_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'COND_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
|
||||
self.refiner_cond_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_COND_MODEL,
|
||||
module_paras.get(
|
||||
'REFINER_COND_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_COND_MODEL') else None
|
||||
self.refiner_diffusion_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_MODEL, module_paras.get(
|
||||
'REFINER_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_MODEL') else None
|
||||
self.tokenizer = TOKENIZERS.build(
|
||||
cfg.MODEL.TOKENIZER,
|
||||
logger=self.logger) if cfg.MODEL.have('TOKENIZER') else None
|
||||
|
||||
if self.tokenizer is not None:
|
||||
self.cond_stage_model['cfg'].KWARGS = {
|
||||
'vocab_size': self.tokenizer.vocab_size
|
||||
}
|
||||
|
||||
def redefine_paras(self, cfg):
|
||||
if cfg.get('PRETRAINED_MODEL', None):
|
||||
assert FS.isfile(cfg.PRETRAINED_MODEL)
|
||||
@@ -159,263 +116,6 @@ class LargenInference():
|
||||
cfg.DIFFUSION_MODEL.RELOAD_MODEL = diffusion_model_path
|
||||
return cfg
|
||||
|
||||
def init_from_modules(self, modules):
|
||||
for k, v in modules.items():
|
||||
self.__setattr__(k, v)
|
||||
|
||||
def infer_model(self, cfg, module_paras=None):
|
||||
module = {
|
||||
'model': None,
|
||||
'cfg': cfg,
|
||||
'device': 'offline',
|
||||
'name': cfg.NAME,
|
||||
'function_info': {},
|
||||
'paras': {}
|
||||
}
|
||||
if module_paras is None:
|
||||
return module
|
||||
function_info = {}
|
||||
paras = {
|
||||
k.lower(): v
|
||||
for k, v in module_paras.get('PARAS', {}).items()
|
||||
}
|
||||
for function in module_paras.get('FUNCTION', []):
|
||||
input_dict = {}
|
||||
for inp in function.get('INPUT', []):
|
||||
if inp.lower() in self.input:
|
||||
input_dict[inp.lower()] = self.input[inp.lower()]
|
||||
function_info[function.NAME] = {
|
||||
'dtype': function.get('DTYPE', 'float32'),
|
||||
'input': input_dict
|
||||
}
|
||||
module['paras'] = paras
|
||||
module['function_info'] = function_info
|
||||
return module
|
||||
|
||||
def init_from_ckpt(self, path, model, 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 = model.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 load(self, module):
|
||||
if module['device'] == 'offline':
|
||||
if module['cfg'].NAME in MODELS.class_map:
|
||||
model = MODELS.build(module['cfg'], logger=self.logger).eval()
|
||||
elif module['cfg'].NAME in BACKBONES.class_map:
|
||||
model = BACKBONES.build(module['cfg'],
|
||||
logger=self.logger).eval()
|
||||
elif module['cfg'].NAME in EMBEDDERS.class_map:
|
||||
model = EMBEDDERS.build(module['cfg'],
|
||||
logger=self.logger).eval()
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if module['cfg'].get('RELOAD_MODEL', None):
|
||||
self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model)
|
||||
module['model'] = model
|
||||
module['device'] = 'cpu'
|
||||
if module['device'] == 'cpu':
|
||||
module['device'] = we.device_id
|
||||
module['model'] = module['model'].to(we.device_id)
|
||||
return module
|
||||
|
||||
def unload(self, module):
|
||||
if module is None:
|
||||
return module
|
||||
mem = get_available_memory()
|
||||
free_mem = int(mem['available'] / (1024**2))
|
||||
total_mem = int(mem['total'] / (1024**2))
|
||||
if free_mem < 0.5 * total_mem:
|
||||
if module['model'] is not None:
|
||||
module['model'] = module['model'].to('cpu')
|
||||
del module['model']
|
||||
module['model'] = None
|
||||
module['device'] = 'offline'
|
||||
print('delete module')
|
||||
else:
|
||||
module['model'] = module['model'].to('cpu')
|
||||
module['device'] = 'cpu'
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return module
|
||||
|
||||
def dynamic_load(self, module=None, name=''):
|
||||
self.logger.info('Loading {} model'.format(name))
|
||||
if name == 'all':
|
||||
for subname in self.loaded_model_name:
|
||||
self.loaded_model[subname] = self.dynamic_load(
|
||||
getattr(self, subname), subname)
|
||||
elif name in self.loaded_model_name:
|
||||
if name in self.loaded_model:
|
||||
if module['cfg'] != self.loaded_model[name]['cfg']:
|
||||
self.unload(self.loaded_model[name])
|
||||
module = self.load(module)
|
||||
self.loaded_model[name] = module
|
||||
return module
|
||||
elif module['device'] == 'cpu' or module['device'] == 'offline':
|
||||
module = self.load(module)
|
||||
return module
|
||||
else:
|
||||
return module
|
||||
else:
|
||||
module = self.load(module)
|
||||
self.loaded_model[name] = module
|
||||
return module
|
||||
else:
|
||||
return self.load(module)
|
||||
|
||||
def dynamic_unload(self, module=None, name='', skip_loaded=False):
|
||||
self.logger.info('Unloading {} model'.format(name))
|
||||
if name == 'all':
|
||||
for name, module in self.loaded_model.items():
|
||||
module = self.unload(self.loaded_model[name])
|
||||
self.loaded_model[name] = module
|
||||
elif name in self.loaded_model_name:
|
||||
if name in self.loaded_model:
|
||||
if not skip_loaded:
|
||||
module = self.unload(self.loaded_model[name])
|
||||
self.loaded_model[name] = module
|
||||
else:
|
||||
self.unload(module)
|
||||
else:
|
||||
self.unload(module)
|
||||
|
||||
def load_default(self, cfg):
|
||||
module_paras = {}
|
||||
if cfg is not None:
|
||||
self.paras = cfg.PARAS
|
||||
self.input = {k.lower(): v for k, v in cfg.INPUT.items()}
|
||||
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
|
||||
module_paras = cfg.MODULES_PARAS
|
||||
return module_paras
|
||||
|
||||
def load_schedule(self, cfg):
|
||||
parameterization = cfg.get('PARAMETERIZATION', 'eps')
|
||||
assert parameterization in [
|
||||
'eps', 'x0', 'v'
|
||||
], 'currently only supporting "eps" and "x0" and "v"'
|
||||
num_timesteps = cfg.get('TIMESTEPS', 1000)
|
||||
|
||||
schedule_args = {
|
||||
k.lower(): v
|
||||
for k, v in cfg.get('SCHEDULE_ARGS', {
|
||||
'NAME': 'logsnr_cosine_interp',
|
||||
'SCALE_MIN': 2.0,
|
||||
'SCALE_MAX': 4.0
|
||||
}).items()
|
||||
}
|
||||
|
||||
zero_terminal_snr = cfg.get('ZERO_TERMINAL_SNR', False)
|
||||
if zero_terminal_snr:
|
||||
assert parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
|
||||
sigmas = noise_schedule(schedule=schedule_args.pop('name'),
|
||||
n=num_timesteps,
|
||||
zero_terminal_snr=zero_terminal_snr,
|
||||
**schedule_args)
|
||||
diffusion = GaussianDiffusion(sigmas=sigmas,
|
||||
prediction_type=parameterization)
|
||||
return diffusion
|
||||
|
||||
def get_batch(self, value_dict, num_samples=1):
|
||||
batch = {}
|
||||
batch_uc = {}
|
||||
N = num_samples
|
||||
device = we.device_id
|
||||
for key in value_dict:
|
||||
if key == 'prompt':
|
||||
if not self.tokenizer:
|
||||
batch['prompt'] = value_dict['prompt']
|
||||
batch_uc['prompt'] = value_dict['negative_prompt']
|
||||
else:
|
||||
batch['tokens'] = self.tokenizer(value_dict['prompt']).to(
|
||||
we.device_id)
|
||||
batch_uc['tokens'] = self.tokenizer(
|
||||
value_dict['negative_prompt']).to(we.device_id)
|
||||
elif key == 'original_size_as_tuple':
|
||||
batch['original_size_as_tuple'] = (torch.tensor(
|
||||
value_dict['original_size_as_tuple']).to(device).repeat(
|
||||
N, 1))
|
||||
elif key == 'crop_coords_top_left':
|
||||
batch['crop_coords_top_left'] = (torch.tensor(
|
||||
value_dict['crop_coords_top_left']).to(device).repeat(
|
||||
N, 1))
|
||||
elif key == 'aesthetic_score':
|
||||
batch['aesthetic_score'] = (torch.tensor(
|
||||
[value_dict['aesthetic_score']]).to(device).repeat(N, 1))
|
||||
batch_uc['aesthetic_score'] = (torch.tensor([
|
||||
value_dict['negative_aesthetic_score']
|
||||
]).to(device).repeat(N, 1))
|
||||
|
||||
elif key == 'target_size_as_tuple':
|
||||
batch['target_size_as_tuple'] = (torch.tensor(
|
||||
value_dict['target_size_as_tuple']).to(device).repeat(
|
||||
N, 1))
|
||||
elif key == 'image':
|
||||
batch[key] = self.load_image(value_dict[key], num_samples=N)
|
||||
else:
|
||||
batch[key] = value_dict[key]
|
||||
|
||||
for key in batch.keys():
|
||||
if key not in batch_uc and isinstance(batch[key], torch.Tensor):
|
||||
batch_uc[key] = torch.clone(batch[key])
|
||||
return batch, batch_uc
|
||||
|
||||
def load_image(self, image, num_samples=1):
|
||||
if isinstance(image, torch.Tensor):
|
||||
pass
|
||||
elif isinstance(image, Image):
|
||||
pass
|
||||
elif isinstance(image, Image):
|
||||
pass
|
||||
|
||||
def get_function_info(self, module, function_name=None):
|
||||
all_function = module['function_info']
|
||||
if function_name in all_function:
|
||||
return function_name, all_function[function_name]['dtype']
|
||||
if function_name is None and len(all_function) == 1:
|
||||
for k, v in all_function.items():
|
||||
return k, v['dtype']
|
||||
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
z = get_model(self.first_stage_model).encode(x)
|
||||
return self.first_stage_model['paras']['scale_factor'] * z
|
||||
|
||||
def decode_first_stage(self, z):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
z = 1. / self.first_stage_model['paras']['scale_factor'] * z
|
||||
return get_model(self.first_stage_model).decode(z)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(self,
|
||||
input,
|
||||
@@ -426,7 +126,10 @@ class LargenInference():
|
||||
cat_uc=True,
|
||||
tuner_model=None,
|
||||
control_model=None,
|
||||
largen_state=False,
|
||||
**kwargs):
|
||||
if not largen_state:
|
||||
raise gr.Error('LARGEN model must be used with LAR-Gen settings')
|
||||
|
||||
value_input = copy.deepcopy(self.input)
|
||||
value_input.update(input)
|
||||
|
||||
@@ -1,24 +1,16 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os.path
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL.Image import Image
|
||||
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.model.network.diffusion.schedules import noise_schedule
|
||||
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
|
||||
TOKENIZERS)
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.utils.env import get_available_memory
|
||||
|
||||
from .control_inference import ControlInference
|
||||
from .diffusion_inference import DiffusionInference
|
||||
from .tuner_inference import TunerInference
|
||||
|
||||
|
||||
@@ -27,7 +19,7 @@ def get_model(model_tuple):
|
||||
return model_tuple['model']
|
||||
|
||||
|
||||
class StyleboothInference():
|
||||
class StyleboothInference(DiffusionInference):
|
||||
'''
|
||||
define vae, unet, text-encoder, tuner, refiner components
|
||||
support to load the components dynamicly.
|
||||
@@ -42,303 +34,6 @@ class StyleboothInference():
|
||||
self.tuner_infer = TunerInference(self.logger)
|
||||
self.control_infer = ControlInference(self.logger)
|
||||
|
||||
def init_from_cfg(self, cfg):
|
||||
self.name = cfg.NAME
|
||||
self.is_default = cfg.get('IS_DEFAULT', False)
|
||||
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
|
||||
assert cfg.have('MODEL')
|
||||
cfg.MODEL = self.redefine_paras(cfg.MODEL)
|
||||
self.diffusion = self.load_schedule(cfg.MODEL.SCHEDULE)
|
||||
self.diffusion_model = self.infer_model(
|
||||
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
|
||||
'DIFFUSION_MODEL',
|
||||
None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None
|
||||
self.first_stage_model = self.infer_model(
|
||||
cfg.MODEL.FIRST_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'FIRST_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None
|
||||
self.cond_stage_model = self.infer_model(
|
||||
cfg.MODEL.COND_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'COND_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
|
||||
self.refiner_cond_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_COND_MODEL,
|
||||
module_paras.get(
|
||||
'REFINER_COND_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_COND_MODEL') else None
|
||||
self.refiner_diffusion_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_MODEL, module_paras.get(
|
||||
'REFINER_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_MODEL') else None
|
||||
self.tokenizer = TOKENIZERS.build(
|
||||
cfg.MODEL.TOKENIZER,
|
||||
logger=self.logger) if cfg.MODEL.have('TOKENIZER') else None
|
||||
|
||||
if self.tokenizer is not None:
|
||||
self.cond_stage_model['cfg'].KWARGS = {
|
||||
'vocab_size': self.tokenizer.vocab_size
|
||||
}
|
||||
|
||||
def redefine_paras(self, cfg):
|
||||
if cfg.get('PRETRAINED_MODEL', None):
|
||||
assert FS.isfile(cfg.PRETRAINED_MODEL)
|
||||
with FS.get_from(cfg.PRETRAINED_MODEL,
|
||||
wait_finish=True) as local_path:
|
||||
if local_path.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(local_path)
|
||||
else:
|
||||
sd = torch.load(local_path, map_location='cpu')
|
||||
first_stage_model_path = os.path.join(
|
||||
os.path.dirname(local_path), 'first_stage_model.pth')
|
||||
cond_stage_model_path = os.path.join(
|
||||
os.path.dirname(local_path), 'cond_stage_model.pth')
|
||||
diffusion_model_path = os.path.join(
|
||||
os.path.dirname(local_path), 'diffusion_model.pth')
|
||||
if (not os.path.exists(first_stage_model_path)
|
||||
or not os.path.exists(cond_stage_model_path)
|
||||
or not os.path.exists(diffusion_model_path)):
|
||||
self.logger.info(
|
||||
'Now read the whole model and rearrange the modules, it may take several mins.'
|
||||
)
|
||||
first_stage_model = OrderedDict()
|
||||
cond_stage_model = OrderedDict()
|
||||
diffusion_model = OrderedDict()
|
||||
for k, v in sd.items():
|
||||
if k.startswith('first_stage_model.'):
|
||||
first_stage_model[k.replace(
|
||||
'first_stage_model.', '')] = v
|
||||
elif k.startswith('conditioner.'):
|
||||
cond_stage_model[k.replace('conditioner.', '')] = v
|
||||
elif k.startswith('cond_stage_model.'):
|
||||
if k.startswith('cond_stage_model.model.'):
|
||||
cond_stage_model[k.replace(
|
||||
'cond_stage_model.model.', '')] = v
|
||||
else:
|
||||
cond_stage_model[k.replace(
|
||||
'cond_stage_model.', '')] = v
|
||||
elif k.startswith('model.diffusion_model.'):
|
||||
diffusion_model[k.replace('model.diffusion_model.',
|
||||
'')] = v
|
||||
else:
|
||||
continue
|
||||
if cfg.have('FIRST_STAGE_MODEL'):
|
||||
with open(first_stage_model_path + 'cache', 'wb') as f:
|
||||
torch.save(first_stage_model, f)
|
||||
os.rename(first_stage_model_path + 'cache',
|
||||
first_stage_model_path)
|
||||
self.logger.info(
|
||||
'First stage model has been processed.')
|
||||
if cfg.have('COND_STAGE_MODEL'):
|
||||
with open(cond_stage_model_path + 'cache', 'wb') as f:
|
||||
torch.save(cond_stage_model, f)
|
||||
os.rename(cond_stage_model_path + 'cache',
|
||||
cond_stage_model_path)
|
||||
self.logger.info(
|
||||
'Cond stage model has been processed.')
|
||||
if cfg.have('DIFFUSION_MODEL'):
|
||||
with open(diffusion_model_path + 'cache', 'wb') as f:
|
||||
torch.save(diffusion_model, f)
|
||||
os.rename(diffusion_model_path + 'cache',
|
||||
diffusion_model_path)
|
||||
self.logger.info('Diffusion model has been processed.')
|
||||
if not cfg.FIRST_STAGE_MODEL.get('PRETRAINED_MODEL', None):
|
||||
cfg.FIRST_STAGE_MODEL.PRETRAINED_MODEL = first_stage_model_path
|
||||
else:
|
||||
cfg.FIRST_STAGE_MODEL.RELOAD_MODEL = first_stage_model_path
|
||||
if not cfg.COND_STAGE_MODEL.get('PRETRAINED_MODEL', None):
|
||||
cfg.COND_STAGE_MODEL.PRETRAINED_MODEL = cond_stage_model_path
|
||||
else:
|
||||
cfg.COND_STAGE_MODEL.RELOAD_MODEL = cond_stage_model_path
|
||||
if not cfg.DIFFUSION_MODEL.get('PRETRAINED_MODEL', None):
|
||||
cfg.DIFFUSION_MODEL.PRETRAINED_MODEL = diffusion_model_path
|
||||
else:
|
||||
cfg.DIFFUSION_MODEL.RELOAD_MODEL = diffusion_model_path
|
||||
return cfg
|
||||
|
||||
def init_from_modules(self, modules):
|
||||
for k, v in modules.items():
|
||||
self.__setattr__(k, v)
|
||||
|
||||
def infer_model(self, cfg, module_paras=None):
|
||||
module = {
|
||||
'model': None,
|
||||
'cfg': cfg,
|
||||
'device': 'offline',
|
||||
'name': cfg.NAME,
|
||||
'function_info': {},
|
||||
'paras': {}
|
||||
}
|
||||
if module_paras is None:
|
||||
return module
|
||||
function_info = {}
|
||||
paras = {
|
||||
k.lower(): v
|
||||
for k, v in module_paras.get('PARAS', {}).items()
|
||||
}
|
||||
for function in module_paras.get('FUNCTION', []):
|
||||
input_dict = {}
|
||||
for inp in function.get('INPUT', []):
|
||||
if inp.lower() in self.input:
|
||||
input_dict[inp.lower()] = self.input[inp.lower()]
|
||||
function_info[function.NAME] = {
|
||||
'dtype': function.get('DTYPE', 'float32'),
|
||||
'input': input_dict
|
||||
}
|
||||
module['paras'] = paras
|
||||
module['function_info'] = function_info
|
||||
return module
|
||||
|
||||
def init_from_ckpt(self, path, model, 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 = model.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 load(self, module):
|
||||
if module['device'] == 'offline':
|
||||
if module['cfg'].NAME in MODELS.class_map:
|
||||
model = MODELS.build(module['cfg'], logger=self.logger).eval()
|
||||
elif module['cfg'].NAME in BACKBONES.class_map:
|
||||
model = BACKBONES.build(module['cfg'],
|
||||
logger=self.logger).eval()
|
||||
elif module['cfg'].NAME in EMBEDDERS.class_map:
|
||||
model = EMBEDDERS.build(module['cfg'],
|
||||
logger=self.logger).eval()
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if module['cfg'].get('RELOAD_MODEL', None):
|
||||
self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model)
|
||||
module['model'] = model
|
||||
module['device'] = 'cpu'
|
||||
if module['device'] == 'cpu':
|
||||
module['device'] = we.device_id
|
||||
module['model'] = module['model'].to(we.device_id)
|
||||
return module
|
||||
|
||||
def unload(self, module):
|
||||
if module is None:
|
||||
return module
|
||||
mem = get_available_memory()
|
||||
free_mem = int(mem['available'] / (1024**2))
|
||||
total_mem = int(mem['total'] / (1024**2))
|
||||
if free_mem < 0.5 * total_mem:
|
||||
if module['model'] is not None:
|
||||
module['model'] = module['model'].to('cpu')
|
||||
del module['model']
|
||||
module['model'] = None
|
||||
module['device'] = 'offline'
|
||||
print('delete module')
|
||||
else:
|
||||
module['model'] = module['model'].to('cpu')
|
||||
module['device'] = 'cpu'
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return module
|
||||
|
||||
def dynamic_load(self, module=None, name=''):
|
||||
self.logger.info('Loading {} model'.format(name))
|
||||
if name == 'all':
|
||||
for subname in self.loaded_model_name:
|
||||
self.loaded_model[subname] = self.dynamic_load(
|
||||
getattr(self, subname), subname)
|
||||
elif name in self.loaded_model_name:
|
||||
if name in self.loaded_model:
|
||||
if module['cfg'] != self.loaded_model[name]['cfg']:
|
||||
self.unload(self.loaded_model[name])
|
||||
module = self.load(module)
|
||||
self.loaded_model[name] = module
|
||||
return module
|
||||
elif module['device'] == 'cpu' or module['device'] == "offline":
|
||||
module = self.load(module)
|
||||
return module
|
||||
else:
|
||||
return module
|
||||
else:
|
||||
module = self.load(module)
|
||||
self.loaded_model[name] = module
|
||||
return module
|
||||
else:
|
||||
return self.load(module)
|
||||
|
||||
def dynamic_unload(self, module=None, name='', skip_loaded=False):
|
||||
self.logger.info('Unloading {} model'.format(name))
|
||||
if name == 'all':
|
||||
for name, module in self.loaded_model.items():
|
||||
module = self.unload(self.loaded_model[name])
|
||||
self.loaded_model[name] = module
|
||||
elif name in self.loaded_model_name:
|
||||
if name in self.loaded_model:
|
||||
if not skip_loaded:
|
||||
module = self.unload(self.loaded_model[name])
|
||||
self.loaded_model[name] = module
|
||||
else:
|
||||
self.unload(module)
|
||||
else:
|
||||
self.unload(module)
|
||||
|
||||
def load_default(self, cfg):
|
||||
module_paras = {}
|
||||
if cfg is not None:
|
||||
self.paras = cfg.PARAS
|
||||
self.input = {k.lower(): v for k, v in cfg.INPUT.items()}
|
||||
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
|
||||
module_paras = cfg.MODULES_PARAS
|
||||
return module_paras
|
||||
|
||||
def load_schedule(self, cfg):
|
||||
parameterization = cfg.get('PARAMETERIZATION', 'eps')
|
||||
assert parameterization in [
|
||||
'eps', 'x0', 'v'
|
||||
], 'currently only supporting "eps" and "x0" and "v"'
|
||||
num_timesteps = cfg.get('TIMESTEPS', 1000)
|
||||
|
||||
schedule_args = {
|
||||
k.lower(): v
|
||||
for k, v in cfg.get('SCHEDULE_ARGS', {
|
||||
'NAME': 'logsnr_cosine_interp',
|
||||
'SCALE_MIN': 2.0,
|
||||
'SCALE_MAX': 4.0
|
||||
}).items()
|
||||
}
|
||||
|
||||
zero_terminal_snr = cfg.get('ZERO_TERMINAL_SNR', False)
|
||||
if zero_terminal_snr:
|
||||
assert parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
|
||||
sigmas = noise_schedule(schedule=schedule_args.pop('name'),
|
||||
n=num_timesteps,
|
||||
zero_terminal_snr=zero_terminal_snr,
|
||||
**schedule_args)
|
||||
diffusion = GaussianDiffusion(sigmas=sigmas,
|
||||
prediction_type=parameterization)
|
||||
return diffusion
|
||||
|
||||
def get_batch(self, value_dict, num_samples=1):
|
||||
batch = {}
|
||||
batch_uc = {}
|
||||
@@ -377,38 +72,6 @@ class StyleboothInference():
|
||||
batch_uc[key] = torch.clone(batch[key])
|
||||
return batch, batch_uc
|
||||
|
||||
def load_image(self, image, num_samples=1):
|
||||
if isinstance(image, torch.Tensor):
|
||||
pass
|
||||
elif isinstance(image, Image):
|
||||
pass
|
||||
elif isinstance(image, Image):
|
||||
pass
|
||||
|
||||
def get_function_info(self, module, function_name=None):
|
||||
all_function = module['function_info']
|
||||
if function_name in all_function:
|
||||
return function_name, all_function[function_name]['dtype']
|
||||
if function_name is None and len(all_function) == 1:
|
||||
for k, v in all_function.items():
|
||||
return k, v['dtype']
|
||||
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
z = get_model(self.first_stage_model).encode(x)
|
||||
return self.first_stage_model['paras']['scale_factor'] * z
|
||||
|
||||
def decode_first_stage(self, z):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
z = 1. / self.first_stage_model['paras']['scale_factor'] * z
|
||||
return get_model(self.first_stage_model).decode(z)
|
||||
|
||||
def encode_condition(self, data, data2=None, type='text'):
|
||||
cond_stage_model = get_model(self.cond_stage_model)
|
||||
assert hasattr(self, 'tokenizer')
|
||||
@@ -457,12 +120,16 @@ class StyleboothInference():
|
||||
cat_uc=True,
|
||||
tuner_model=None,
|
||||
control_model=None,
|
||||
stylebooth_state=False,
|
||||
style_edit_image=None,
|
||||
style_exemplar_image=None,
|
||||
style_guide_scale_text=None,
|
||||
style_guide_scale_image=None,
|
||||
**kwargs):
|
||||
|
||||
if not stylebooth_state:
|
||||
raise gr.Error('EDIT model must be used with StyleBooth settings')
|
||||
|
||||
value_input = copy.deepcopy(self.input)
|
||||
value_input.update(input)
|
||||
print(value_input)
|
||||
|
||||
@@ -68,6 +68,7 @@ class CheckpointHook(Hook):
|
||||
self.save_best_by = cfg.get('SAVE_BEST_BY', '')
|
||||
self.push_to_hub = cfg.get('PUSH_TO_HUB', False)
|
||||
self.hub_model_id = cfg.get('HUB_MODEL_ID', None)
|
||||
self.hub_private = cfg.get('HUB_PRIVATE', False)
|
||||
self.disable_save_snapshot = cfg.get('DISABLE_SNAPSHOT', False)
|
||||
self.last_ckpt = None
|
||||
if self.save_best and not self.save_best_by:
|
||||
@@ -254,7 +255,7 @@ class CheckpointHook(Hook):
|
||||
|
||||
push_to_hub(repo_name=self.hub_model_id,
|
||||
output_dir=self.last_ckpt,
|
||||
private=False)
|
||||
private=self.hub_private)
|
||||
current_dir = os.path.dirname(__file__)
|
||||
base_path = os.sep.join(current_dir.split(os.sep)[:-4])
|
||||
base_path = os.path.join(base_path, 'config')
|
||||
|
||||
@@ -242,7 +242,7 @@ class InferenceUI():
|
||||
self.check_box_for_setting, self.setting_tab, *gr_states,
|
||||
*gr_tabs, self.model_manage_ui.diffusion_model
|
||||
],
|
||||
queue=False)
|
||||
queue=True)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
@@ -198,6 +198,17 @@ class GalleryUI(UIBase):
|
||||
'seed': int(args.pop('image_seed')),
|
||||
})
|
||||
|
||||
def load_init(args):
|
||||
cur_pipe_name = self.pipe_manager.model_level_info[
|
||||
args['diffusion_model']]['pipeline'][0]
|
||||
for sub_name, sub_pipe in self.pipe_manager.pipeline_level_modules.items(
|
||||
):
|
||||
if sub_name == cur_pipe_name:
|
||||
continue
|
||||
if len(sub_pipe.loaded_model) > 0:
|
||||
sub_pipe.dynamic_unload(name='all')
|
||||
print(f'Unloading {sub_name} modules')
|
||||
|
||||
args = dict(zip(self.component_mapping.keys(), args))
|
||||
largen_history = args.pop('largen_history')
|
||||
|
||||
@@ -206,7 +217,7 @@ class GalleryUI(UIBase):
|
||||
tuner_init(args)
|
||||
input_init(args)
|
||||
appedix_init(args)
|
||||
|
||||
load_init(args)
|
||||
results = current_pipeline(**args)
|
||||
|
||||
images = []
|
||||
|
||||
@@ -115,6 +115,8 @@ class StyleboothUI(UIBase):
|
||||
# cache_examples=True)
|
||||
gallery_ui = kwargs.pop('gallery_ui', None)
|
||||
gallery_ui.register_components({
|
||||
'stylebooth_state':
|
||||
self.state,
|
||||
'style_edit_image':
|
||||
self.edit_image,
|
||||
'style_exemplar_image':
|
||||
|
||||
@@ -794,7 +794,9 @@ class DatasetGalleryUI(UIBase):
|
||||
image = upload_image['image']
|
||||
else:
|
||||
image = upload_image
|
||||
target_image = processor_ins(image.convert('RGB'))
|
||||
target_image = processor_ins(image.convert('RGB'),
|
||||
height_ratio=height_ratio,
|
||||
width_ratio=width_ratio)
|
||||
w, h = target_image.size
|
||||
ret_image_gallery = gr.Gallery()
|
||||
ret_upload_image = gr.Image(target_image)
|
||||
|
||||
@@ -521,6 +521,14 @@ class Text2ImageDataCard(BaseDataCard):
|
||||
except Exception:
|
||||
pass
|
||||
file_list = self.load_train_file(new_train_list, data_folder)
|
||||
# remove unused data
|
||||
for one_dir in FS.walk_dir(local_dataset_folder):
|
||||
if "images" in one_dir or one_dir.endswith("file.csv") or one_dir.endswith("train.csv"):
|
||||
continue
|
||||
try:
|
||||
os.system(f"rm -rf {one_dir}")
|
||||
except:
|
||||
pass
|
||||
return file_list
|
||||
|
||||
def load_train_file(self, file_path, data_folder):
|
||||
|
||||
@@ -17,12 +17,18 @@ class TaskStatus():
|
||||
def __init__(self):
|
||||
self.pid = -1
|
||||
self.retcode = -999
|
||||
self.error_msg = None
|
||||
self.out_msg = None
|
||||
self.error_log = None
|
||||
|
||||
def __repr__(self):
|
||||
return f'Process {self.pid} retcode {self.retcode}, error msg: {self.error_msg}.'
|
||||
|
||||
@property
|
||||
def error_msg(self):
|
||||
if os.path.exists(self.error_log):
|
||||
return "\n".join(open(self.error_log, "r").readlines()[-50:])
|
||||
else:
|
||||
return "No error msg."
|
||||
|
||||
|
||||
def kill_job(pid):
|
||||
try:
|
||||
@@ -46,8 +52,11 @@ class Trainer():
|
||||
|
||||
def __call__(self, task_name):
|
||||
torch.cuda.empty_cache()
|
||||
error_folder = "./error_logs"
|
||||
os.makedirs(error_folder, exist_ok=True)
|
||||
self.status_message.error_log = f"{error_folder}/{int(time.time())}.log"
|
||||
cmd = f'PYTHONPATH=. python {self.run_script} ' \
|
||||
f'--cfg={task_name}/train.yaml'
|
||||
f'--cfg={task_name}/train.yaml 2> {self.status_message.error_log}'
|
||||
# cmd = [f"python {self.run_script}"]
|
||||
print(cmd)
|
||||
try:
|
||||
@@ -57,11 +66,9 @@ class Trainer():
|
||||
self.status_message.retcode = self.proc.wait(
|
||||
) # self.proc.wait(3600*24*14)
|
||||
# self.status_message.error_msg = self.proc.stderr.read()
|
||||
self.status_message.error_msg = ''
|
||||
except Exception:
|
||||
self.status_message.retcode = -2
|
||||
# self.status_message.error_msg = self.proc.stderr.read()
|
||||
self.status_message.error_msg = ''
|
||||
|
||||
def terminate(self):
|
||||
print(f'Terminate {self.proc.pid} ...')
|
||||
@@ -118,7 +125,7 @@ class TrainManager():
|
||||
err_msg = now_task['train_status'].error_msg
|
||||
message = f'''
|
||||
Training failed! \n
|
||||
Error msg: {err_msg[-1000:]} \n
|
||||
Error msg: {err_msg} \n
|
||||
Take time [ {duration:.4f}s ] \n
|
||||
{self.check_memory()}
|
||||
'''
|
||||
|
||||
@@ -451,7 +451,7 @@ class ModelUI(UIBase):
|
||||
# manager.inference.tuner_ui.tuner_desc,
|
||||
# manager.inference.tuner_ui.tuner_prompt_example
|
||||
],
|
||||
queue=False)
|
||||
queue=True)
|
||||
|
||||
def export_train_log(model_name):
|
||||
current_log_folder = os.path.join(self.work_dir, model_name)
|
||||
|
||||
@@ -187,7 +187,7 @@ class InfoUI(UIBase):
|
||||
manager.inference.model_manage_ui.diffusion_model,
|
||||
manager.inference.tuner_ui.custom_tuner_model
|
||||
],
|
||||
queue=False)
|
||||
queue=True)
|
||||
|
||||
def export_zip(tuner_name, base_model, tuner_type):
|
||||
sub_dir = f'{base_model}-{tuner_type}'
|
||||
|
||||
Reference in New Issue
Block a user