update v1.2.0

This commit is contained in:
LouieStark
2024-11-01 10:15:27 +08:00
parent eac03e9856
commit e8d8e63cba
71 changed files with 5926 additions and 991 deletions
+75 -57
View File
@@ -10,23 +10,24 @@ from functools import partial
import numpy as np
import torch
import torch.nn as nn
from torch.distributed.fsdp import FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import (MixedPrecision, ShardingStrategy,
StateDictType)
from torch.distributed.fsdp.wrap import lambda_auto_wrap_policy
from torch.nn.parallel import DistributedDataParallel
from tqdm import tqdm
from scepter.modules.data.dataset import DATASETS
from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS
from scepter.modules.opt.optimizers import OPTIMIZERS
from scepter.modules.solver import BaseSolver
from scepter.modules.solver.registry import SOLVERS
from scepter.modules.utils.config import Config, dict_to_yaml
from scepter.modules.utils.data import transfer_data_to_cuda
from scepter.modules.utils.distribute import we
from scepter.modules.utils.probe import ProbeData
from torch.distributed.fsdp import FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import (MixedPrecision, ShardingStrategy,
StateDictType)
from torch.distributed.fsdp.wrap import (lambda_auto_wrap_policy,
size_based_auto_wrap_policy)
from torch.nn.parallel import DistributedDataParallel
from tqdm import tqdm
from .base_solver import BaseSolver
sharding_strategy_map = {
'full_shard': ShardingStrategy.FULL_SHARD,
@@ -40,13 +41,14 @@ def shard_model(model,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
buffer_dtype=torch.float32,
fsdp_group = ['blocks'],
fsdp_group=['blocks'],
sharding_strategy=ShardingStrategy.FULL_SHARD,
sync_module_states=False):
wrap_modules = []
for module_name in fsdp_group:
if hasattr(model, module_name):
if isinstance(getattr(model, module_name), (list, tuple, nn.ModuleList)):
if isinstance(getattr(model, module_name),
(list, tuple, nn.ModuleList)):
wrap_modules.extend([m for m in getattr(model, module_name)])
else:
wrap_modules.extend([getattr(model, module_name)])
@@ -209,7 +211,7 @@ class LatentDiffusionSolver(BaseSolver):
self.sample_args = cfg.get('SAMPLE_ARGS', None)
self.tuner_cfg = cfg.get('TUNER', None)
self.freeze_cfg = cfg.get('FREEZE', None)
self.log_train_num = cfg.get("LOG_TRAIN_NUM", -1)
self.log_train_num = cfg.get('LOG_TRAIN_NUM', -1)
def set_up(self):
self.construct_data()
@@ -285,11 +287,10 @@ class LatentDiffusionSolver(BaseSolver):
from fairscale.nn.data_parallel import ShardedDataParallel
from fairscale.optim.oss import OSS
if hasattr(self.model, 'ignored_parameters'):
train_params, ignored_params = self.model.parameters(
train_params, _ = self.model.parameters(
), self.model.ignored_parameters()
else:
train_params, ignored_params = self.model.parameters(
), None
train_params, _ = self.model.parameters(), None
self.optimizer = OSS(params=train_params,
optim=torch.optim.AdamW,
lr=self.cfg.OPTIMIZER.LEARNING_RATE)
@@ -302,16 +303,18 @@ class LatentDiffusionSolver(BaseSolver):
sub_module = get_module(self.model, module)
if sub_module is not None:
sub_module = shard_model(
sub_module,
device_id=we.device_id,
param_dtype=self.dtype,
reduce_dtype=self.reduce_dtype,
buffer_dtype=self.buffer_dtype,
sharding_strategy=sharding_strategy_map[self.model_shard],
sync_module_states=True)
sub_module,
device_id=we.device_id,
param_dtype=self.dtype,
reduce_dtype=self.reduce_dtype,
buffer_dtype=self.buffer_dtype,
sharding_strategy=sharding_strategy_map[
self.model_shard],
sync_module_states=True)
set_module(self.model, module, sub_module)
elif isinstance(module, (dict, Config)):
sub_module = get_module(self.model, module["MODULE"])
sub_module = get_module(self.model,
module['MODULE'])
if sub_module is not None:
sub_module = shard_model(
sub_module,
@@ -319,10 +322,13 @@ class LatentDiffusionSolver(BaseSolver):
param_dtype=self.dtype,
reduce_dtype=self.reduce_dtype,
buffer_dtype=self.buffer_dtype,
fsdp_group=module.get("FSDP_GROUP", ["blocks"]),
sharding_strategy=sharding_strategy_map[self.model_shard],
fsdp_group=module.get(
'FSDP_GROUP', ['blocks']),
sharding_strategy=sharding_strategy_map[
self.model_shard],
sync_module_states=True)
set_module(self.model, module["MODULE"], sub_module)
set_module(self.model, module['MODULE'],
sub_module)
else:
self.logger.warning(
'FSDP_SHARD_MODULES is None, which means wraping the whold model as the '
@@ -390,6 +396,7 @@ class LatentDiffusionSolver(BaseSolver):
else:
self.scaler = None
self.logger.info(self.model)
def load_checkpoint(self, checkpoint: dict):
"""
Load checkpoint function
@@ -443,10 +450,10 @@ class LatentDiffusionSolver(BaseSolver):
f'Load checkpoint for optimizer {module}.')
else:
self.optimizer.load_state_dict(checkpoint['optimizer'])
self.logger.info(f'Load checkpoint for optimizer.')
self.logger.info('Load checkpoint for optimizer.')
if 'scaler' in checkpoint and self.scaler:
self.scaler.load_state_dict(checkpoint['scaler'])
self.logger.info(f'Load checkpoint for scaler.')
self.logger.info('Load checkpoint for scaler.')
self.logger.info('Load checkpoint finished.')
def save_checkpoint(self) -> dict:
@@ -498,7 +505,7 @@ class LatentDiffusionSolver(BaseSolver):
model = self.model
if self.save_modules is not None:
for module in self.save_modules:
current_module = get_module(self.model, module)
current_module = get_module(model, module)
if current_module is not None:
ckpt['model'][module] = current_module.state_dict()
else:
@@ -510,7 +517,7 @@ class LatentDiffusionSolver(BaseSolver):
model = self.model
if self.save_modules is not None:
for module in self.save_modules:
current_module = get_module(self.model, module)
current_module = get_module(model, module)
if current_module is not None:
ckpt['model'][module] = current_module.state_dict()
else:
@@ -624,12 +631,12 @@ class LatentDiffusionSolver(BaseSolver):
# the inference image use
ret_images, ret_labels = [], []
if 'hint' in result:
ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(f"Control Image")
ret_images.append(
(result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_images.append(
(result['hint'][:result['image'].shape[0]].permute(
1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append('Control Image')
ret_images.append((result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(result['prompt'] +
" <font color='red'> |NegPrompt| </font> " +
result['n_prompt'])
@@ -671,15 +678,15 @@ class LatentDiffusionSolver(BaseSolver):
# the inference image use
ret_images, ret_labels = [], []
if 'hint' in result:
ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(f"Control Image")
ret_images.append(
(result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_images.append(
(result['hint'][:result['image'].shape[0]].permute(
1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append('Control Image')
ret_images.append((result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(result['prompt'] +
" <font color='red'> |NegPrompt| </font> " +
result['n_prompt'])
" <font color='red'> |NegPrompt| </font> " +
result['n_prompt'])
log_data.append(ret_images)
log_label.append(ret_labels)
ori_label.append(result['prompt'])
@@ -710,7 +717,9 @@ class LatentDiffusionSolver(BaseSolver):
from swift import Swift
model = Swift.prepare_model(self.model, config=swift_cfg_dict)
self.logger.info([(key, param.shape) for key, param in model.named_parameters() if param.requires_grad])
self.logger.info([(key, param.shape)
for key, param in model.named_parameters()
if param.requires_grad])
return model
def freeze(self, freeze_cfg, model=None):
@@ -772,7 +781,9 @@ class LatentDiffusionSolver(BaseSolver):
for name, param in freeze_model.named_parameters():
if re.match(train_part, name):
param.requires_grad = True
self.logger.info([(key, param.shape) for key, param in freeze_model.named_parameters() if param.requires_grad])
self.logger.info([(key, param.shape)
for key, param in freeze_model.named_parameters()
if param.requires_grad])
return model
@torch.no_grad()
@@ -815,30 +826,37 @@ class LatentDiffusionSolver(BaseSolver):
@property
def probe_data(self):
if not we.debug and self.mode == 'train':
batch_data = transfer_data_to_cuda(self.current_batch_data[self.mode])
batch_data = transfer_data_to_cuda(
self.current_batch_data[self.mode])
self.eval_mode()
with torch.autocast(device_type='cuda',
enabled=self.use_amp,
dtype=self.dtype):
batch_data['log_num'] = self.log_train_num
results = self.run_step_eval(batch_data)
images = batch_data['image'] if 'image' in batch_data else [None] * len(results)
images = batch_data['image'] if 'image' in batch_data else [
None
] * len(results)
self.train_mode()
log_data, log_label = [], []
for result, image in zip(results, images):
ret_images, ret_labels = [], []
if 'hint' in result:
ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_images.append(
(result['hint'][:result['image'].shape[0]].permute(
1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
if image is not None:
image = torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0)
ret_images.append((image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append(f'target image')
ret_images.append((image.permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append('target image')
ret_images.append((result['image'].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append(result['prompt']
+ " <font color='red'> |NegPrompt| </font> "
+ result['n_prompt'])
ret_images.append(
(result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(result['prompt'] +
" <font color='red'> |NegPrompt| </font> " +
result['n_prompt'])
log_data.append(ret_images)
log_label.append(ret_labels)
self.register_probe({
@@ -929,4 +947,4 @@ class LatentDiffusionSolver(BaseSolver):
logger.info(
f'Load ema frozen params {ema_param_numel} / {all_param_numel} = '
f'{ema_param_numel / all_param_numel:.2%}, '
f'frozen part: {ema_param_dict}.')
f'frozen part: {ema_param_dict}.')