update v1.2.0
This commit is contained in:
@@ -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}.')
|
||||
|
||||
Reference in New Issue
Block a user