update v1.1.0
This commit is contained in:
@@ -18,7 +18,10 @@ from scepter.modules.solver.hooks import HOOKS
|
||||
from scepter.modules.utils.config import Config, dict_to_yaml
|
||||
from scepter.modules.utils.data import transfer_data_to_cuda
|
||||
from scepter.modules.utils.directory import get_relative_folder, osp_path
|
||||
from scepter.modules.utils.distribute import dist, gather_data, we
|
||||
from scepter.modules.utils.distribute import (
|
||||
dist, gather_data, we, all_reduce,
|
||||
_serialize_to_tensor, broadcast, _unserialize_from_tensor,
|
||||
all_reduce, barrier)
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.logger import get_logger, init_logger
|
||||
from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe,
|
||||
@@ -184,6 +187,20 @@ try:
|
||||
except Exception as e:
|
||||
warnings.warn(f'{e}')
|
||||
|
||||
def async_str(text):
|
||||
broadcast_size = torch.zeros(1, dtype=torch.long).to(we.device_id)
|
||||
if we.rank == 0:
|
||||
text_tensor = _serialize_to_tensor(text).to(we.device_id)
|
||||
broadcast_size[0] = len(text_tensor)
|
||||
broadcast(broadcast_size, src=0)
|
||||
broadcast(text_tensor, src=0)
|
||||
else:
|
||||
broadcast(broadcast_size, src=0)
|
||||
text_tensor = torch.empty((broadcast_size[0],), dtype=torch.uint8).to(we.device_id)
|
||||
broadcast(text_tensor, src=0)
|
||||
text = _unserialize_from_tensor(text_tensor)
|
||||
return text
|
||||
|
||||
|
||||
class BaseSolver(object, metaclass=ABCMeta):
|
||||
""" Base Solver.
|
||||
@@ -198,18 +215,12 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
'description': 'The precision for train process.'
|
||||
},
|
||||
'FILE_SYSTEM': {},
|
||||
'ACCU_STEP': {
|
||||
'value':
|
||||
1,
|
||||
'description':
|
||||
'When use ddp, the grad accumulate steps for each process.'
|
||||
},
|
||||
'RESUME_FROM': {
|
||||
'value': '',
|
||||
'description': 'Resume from some state of training!'
|
||||
},
|
||||
'MAX_EPOCHS': {
|
||||
'value': 10,
|
||||
'value': -1,
|
||||
'description': 'Max epochs for training.'
|
||||
},
|
||||
'NUM_FOLDS': {
|
||||
@@ -252,31 +263,31 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
# initialize some hyperparameters
|
||||
self.cfg = cfg
|
||||
self.logger = logger
|
||||
self.file_system = cfg.get('FILE_SYSTEM', None)
|
||||
self.work_dir: str = cfg.WORK_DIR
|
||||
self.work_dir: str = async_str(cfg.WORK_DIR)
|
||||
barrier()
|
||||
self.pl_dir = self.work_dir
|
||||
self.log_file = osp_path(self.work_dir, cfg.LOG_FILE)
|
||||
self.optimizer, self.lr_scheduler = None, None
|
||||
self.cfg = cfg
|
||||
self.logger = logger
|
||||
self.resume_from: str = cfg.RESUME_FROM
|
||||
self.max_epochs: int = cfg.MAX_EPOCHS
|
||||
self.resume_from: str = cfg.get("RESUME_FROM", None)
|
||||
self.max_epochs: int = cfg.get("MAX_EPOCHS", -1)
|
||||
self.use_pl = we.use_pl
|
||||
self.train_precision = self.cfg.get('TRAIN_PRECISION', 32)
|
||||
self._mode_set = set()
|
||||
self._mode = 'train'
|
||||
self.probe_ins = {}
|
||||
self.collect_probe_ins = {}
|
||||
self.clear_probe_ins = {}
|
||||
self._num_folds: int = 1
|
||||
if not self.use_pl:
|
||||
world_size = we.world_size
|
||||
if world_size > 1:
|
||||
self._num_folds: int = cfg.NUM_FOLDS
|
||||
self._num_folds: int = cfg.get("NUM_FOLDS", 1)
|
||||
if cfg.have('MODE'):
|
||||
self._mode_set.add(cfg.MODE)
|
||||
self._mode = cfg.MODE
|
||||
if we.is_distributed:
|
||||
self.accu_step = cfg.get('ACCU_STEP', 1)
|
||||
|
||||
self.do_step = True
|
||||
self.hooks_dict = {'train': [], 'eval': [], 'test': []}
|
||||
@@ -305,6 +316,8 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
self._prefix = FS.get_fs_client(self.work_dir).get_prefix()
|
||||
if not FS.exists(self.work_dir):
|
||||
FS.make_dir(self.work_dir)
|
||||
assert self.cfg.have('MODEL')
|
||||
self.cfg.MODEL.WORK_DIR = self.work_dir
|
||||
self.logger.info(
|
||||
f"Parse work dir {self.work_dir}'s prefix is {self._prefix}")
|
||||
|
||||
@@ -325,6 +338,7 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
def __setattr__(self, key, value):
|
||||
if isinstance(value, BaseModel):
|
||||
self.probe_ins[key] = value.probe_data
|
||||
self.collect_probe_ins[key] = value.collect_probe
|
||||
self.clear_probe_ins[key] = value.clear_probe
|
||||
super().__setattr__(key, value)
|
||||
|
||||
@@ -666,7 +680,7 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
return self._iter[self._mode]
|
||||
|
||||
@property
|
||||
def probe_data(self):
|
||||
def probe_data_dict(self):
|
||||
return self._probe_data[self._mode]
|
||||
|
||||
@property
|
||||
@@ -736,8 +750,32 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
else:
|
||||
self._dist_data[self.mode][key][k] = v
|
||||
|
||||
@property
|
||||
def collect_probe(self):
|
||||
probe_data_dict = self._probe_data[self.mode]
|
||||
for k, func in self.collect_probe_ins.items():
|
||||
for kk, vv in func().items():
|
||||
probe_data_dict[f'{k}/{kk}'] = vv
|
||||
return probe_data_dict
|
||||
|
||||
@property
|
||||
def probe_data(self): # noqa
|
||||
if hasattr(self, f'{self.mode}_pre_save_paras'):
|
||||
pre_save_paras = getattr(self, f'{self.mode}_pre_save_paras')
|
||||
save_folder = pre_save_paras['save_folder']
|
||||
save_probe_prefix = pre_save_paras['save_probe_prefix']
|
||||
step = pre_save_paras['step']
|
||||
save_image_postfix = pre_save_paras.get('save_image_postfix', 'jpg')
|
||||
save_video_postfix = pre_save_paras.get('save_video_postfix', 'mp4')
|
||||
for k, v in self.collect_probe.items():
|
||||
if save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder, save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(save_folder, k.replace('/', '_') + f'_step_{step}')
|
||||
v.presave(prefix = ret_prefix,
|
||||
image_postfix = save_image_postfix,
|
||||
video_postfix = save_video_postfix,
|
||||
rank = we.rank)
|
||||
gather_probe_data = gather_data(self._probe_data[self.mode])
|
||||
_dist_data_list = gather_data([self._dist_data[self.mode] or {}])
|
||||
if not we.rank == 0:
|
||||
@@ -836,9 +874,10 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
for key in keys:
|
||||
value = data_dict[key]
|
||||
if isinstance(value, torch.Tensor) and value.ndim == 0:
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
if we.is_distributed:
|
||||
value = value.data.clone()
|
||||
dist.all_reduce(value.div_(dist.get_world_size()))
|
||||
all_reduce(value, group=we.data_parallel_group)
|
||||
value = value/we.data_group_world_size
|
||||
ret[key] = value
|
||||
else:
|
||||
ret[key] = value
|
||||
@@ -911,7 +950,7 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('solvername',
|
||||
return dict_to_yaml('SOLVER',
|
||||
__class__.__name__,
|
||||
BaseSolver.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -2,11 +2,14 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os
|
||||
import re
|
||||
import warnings
|
||||
from collections import OrderedDict, defaultdict
|
||||
from functools import partial
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.cuda.amp as amp
|
||||
import torch.nn as nn
|
||||
from scepter.modules.data.dataset import DATASETS
|
||||
from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS
|
||||
from scepter.modules.opt.optimizers import OPTIMIZERS
|
||||
@@ -16,19 +19,80 @@ 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 (BackwardPrefetch, CPUOffload,
|
||||
FullStateDictConfig,
|
||||
FullyShardedDataParallel, MixedPrecision,
|
||||
ShardingStrategy, StateDictType)
|
||||
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
|
||||
|
||||
sharding_strategy_map = {
|
||||
'full_shard': ShardingStrategy.FULL_SHARD,
|
||||
'shard_grad_op': ShardingStrategy.SHARD_GRAD_OP
|
||||
'shard_grad_op': ShardingStrategy.SHARD_GRAD_OP,
|
||||
'hybrid_shard': ShardingStrategy.HYBRID_SHARD
|
||||
}
|
||||
|
||||
|
||||
def shard_model(model,
|
||||
device_id,
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
buffer_dtype=torch.float32,
|
||||
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)):
|
||||
wrap_modules.extend([m for m in getattr(model, module_name)])
|
||||
else:
|
||||
wrap_modules.extend([getattr(model, module_name)])
|
||||
else:
|
||||
warnings.warn("Can't find module {} in model".format(module_name))
|
||||
return FSDP(
|
||||
module=model,
|
||||
process_group=None,
|
||||
sharding_strategy=sharding_strategy,
|
||||
auto_wrap_policy=partial(
|
||||
# size_based_auto_wrap_policy, min_num_params=int(1e6),
|
||||
lambda_auto_wrap_policy,
|
||||
lambda_fn=lambda m: m in wrap_modules),
|
||||
mixed_precision=MixedPrecision(param_dtype=param_dtype,
|
||||
reduce_dtype=reduce_dtype,
|
||||
buffer_dtype=buffer_dtype),
|
||||
device_id=device_id,
|
||||
sync_module_states=sync_module_states)
|
||||
|
||||
|
||||
def get_module(instance, sub_module):
|
||||
sub_module_list = sub_module.split('.')
|
||||
for sub_mod in sub_module_list:
|
||||
if sub_mod == '':
|
||||
continue
|
||||
if hasattr(instance, sub_mod):
|
||||
instance = getattr(instance, sub_mod)
|
||||
else:
|
||||
return None
|
||||
return instance
|
||||
|
||||
|
||||
def set_module(instance, sub_module, value):
|
||||
sub_module_list = sub_module.split('.')
|
||||
instance_list = []
|
||||
for sub_mod in sub_module_list:
|
||||
if hasattr(instance, sub_mod):
|
||||
instance_list.append((instance, sub_mod))
|
||||
instance = getattr(instance, sub_mod)
|
||||
instance = value
|
||||
if len(instance_list) > 0:
|
||||
for parents_instance, sub_mod in instance_list[::-1]:
|
||||
setattr(parents_instance, sub_mod, instance)
|
||||
instance = parents_instance
|
||||
|
||||
|
||||
@SOLVERS.register_class()
|
||||
class LatentDiffusionSolver(BaseSolver):
|
||||
para_dict = {
|
||||
@@ -61,6 +125,24 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
'description':
|
||||
f'The shard strategy for fsdp, select from {list(sharding_strategy_map.keys())}',
|
||||
},
|
||||
'FSDP_REDUCE_DTYPE': {
|
||||
'value': 'float32',
|
||||
'description': 'The dtype of reduce in FSDP.'
|
||||
},
|
||||
'FSDP_BUFFER_DTYPE': {
|
||||
'value': 'float32',
|
||||
'description': 'The dtype of buffer in FSDP.'
|
||||
},
|
||||
'FSDP_SHARD_MODULES': {
|
||||
'value': ['model'],
|
||||
'description': 'The modules to be sharded in FSDP.'
|
||||
},
|
||||
'SAVE_MODULES': {
|
||||
'value':
|
||||
None,
|
||||
'description':
|
||||
'The modules to be saved, default is None to save all modules in checkpoint file.'
|
||||
},
|
||||
'IMAGE_LOG_STEP': {
|
||||
'value': 2000,
|
||||
'description': 'The interval for image log.',
|
||||
@@ -112,6 +194,13 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
else:
|
||||
self.logger.info('Use default backend.')
|
||||
self.model_shard = cfg.get('SHARDING_STRATEGY', 'full_shard')
|
||||
self.reduce_dtype = getattr(torch,
|
||||
cfg.get('FSDP_REDUCE_DTYPE', 'float32'))
|
||||
self.buffer_dtype = getattr(torch,
|
||||
cfg.get('FSDP_BUFFER_DTYPE', 'float32'))
|
||||
self.shard_modules = cfg.get('FSDP_SHARD_MODULES', ['model'])
|
||||
self.save_modules = cfg.get('SAVE_MODULES', ['model'])
|
||||
self.train_modules = cfg.get('TRAIN_MODULES', ['model'])
|
||||
self.image_log_step = cfg.get('IMAGE_LOG_STEP', 2000)
|
||||
self._image_out = defaultdict(list)
|
||||
self.load_model_only = cfg.get('LOAD_MODEL_ONLY', False)
|
||||
@@ -120,6 +209,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)
|
||||
|
||||
def set_up(self):
|
||||
self.construct_data()
|
||||
@@ -178,7 +268,7 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.model = self.model.to(we.device_id)
|
||||
|
||||
def init_lr(self):
|
||||
rescale_lr = self.cfg.get('RESCALE_LR', True)
|
||||
rescale_lr = self.cfg.get('RESCALE_LR', False)
|
||||
if rescale_lr and 'train' in self.datas and self.cfg.have('OPTIMIZER'):
|
||||
if we.world_size > 1:
|
||||
all_batch_size = self.datas['train'].batch_size * we.world_size
|
||||
@@ -188,32 +278,76 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.cfg.OPTIMIZER.LEARNING_RATE /= 640
|
||||
|
||||
def init_opti(self):
|
||||
if hasattr(self.model, 'ignored_parameters'):
|
||||
train_params, ignored_params = self.model.parameters(
|
||||
), self.model.ignored_parameters()
|
||||
else:
|
||||
train_params, ignored_params = self.model.parameters(), None
|
||||
import torch.cuda.amp as amp
|
||||
|
||||
if we.is_distributed:
|
||||
if self.use_fairscale:
|
||||
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(
|
||||
), self.model.ignored_parameters()
|
||||
else:
|
||||
train_params, ignored_params = self.model.parameters(
|
||||
), None
|
||||
self.optimizer = OSS(params=train_params,
|
||||
optim=torch.optim.AdamW,
|
||||
lr=self.cfg.OPTIMIZER.LEARNING_RATE)
|
||||
self.model = ShardedDataParallel(self.model, self.optimizer)
|
||||
elif self.use_fsdp:
|
||||
mixed_precision = MixedPrecision(param_dtype=self.dtype,
|
||||
reduce_dtype=self.dtype,
|
||||
buffer_dtype=self.dtype)
|
||||
sharding_strategy = sharding_strategy_map[self.model_shard]
|
||||
self.model = FullyShardedDataParallel(
|
||||
self.model,
|
||||
mixed_precision=mixed_precision,
|
||||
cpu_offload=CPUOffload(offload_params=False),
|
||||
sharding_strategy=sharding_strategy,
|
||||
backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
|
||||
device_id=torch.cuda.current_device(),
|
||||
ignored_parameters=ignored_params)
|
||||
shard_fn = partial
|
||||
if self.shard_modules is not None:
|
||||
for module in self.shard_modules:
|
||||
if isinstance(module, str):
|
||||
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)
|
||||
set_module(self.model, module, sub_module)
|
||||
elif isinstance(module, (dict, Config)):
|
||||
sub_module = get_module(self.model, module["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,
|
||||
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)
|
||||
else:
|
||||
self.logger.warning(
|
||||
'FSDP_SHARD_MODULES is None, which means wraping the whold model as the '
|
||||
'fsdp instance. When using FSDP, it is necessary to specify the modules '
|
||||
'to be wrapped; otherwise, there may be a situation where submodules are '
|
||||
'not the root module, which can lead to unexpected issues. Specify the '
|
||||
'modules to be wrapped by setting FSDP_SHARD_MODULES to a list of modules '
|
||||
'that need wrapping.')
|
||||
self.model = shard_fn(self.model)
|
||||
train_params = []
|
||||
if self.train_modules is None:
|
||||
self.logger.warning(
|
||||
'When using FSDP, it is necessary to explicitly specify the modules to be '
|
||||
'trained or the modules for which gradients will be computed, otherwise, '
|
||||
'there will be issues with gradient calculation.')
|
||||
assert self.train_modules is None
|
||||
else:
|
||||
self.logger.info(
|
||||
f"The modules {','.join(self.train_modules)} 's parameters will be backwarded."
|
||||
)
|
||||
for module in self.train_modules:
|
||||
if hasattr(self.model, module):
|
||||
current_module = getattr(self.model, module)
|
||||
train_params += list(current_module.parameters())
|
||||
|
||||
self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=train_params)
|
||||
@@ -223,14 +357,16 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.model,
|
||||
device_ids=[torch.cuda.current_device()],
|
||||
output_device=torch.cuda.current_device(),
|
||||
find_unused_parameters=True)
|
||||
self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=train_params)
|
||||
find_unused_parameters=False)
|
||||
self.optimizer = OPTIMIZERS.build(
|
||||
self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=self.model.parameters())
|
||||
else:
|
||||
self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=train_params)
|
||||
self.optimizer = OPTIMIZERS.build(
|
||||
self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=self.model.parameters())
|
||||
|
||||
if self.cfg.have('LR_SCHEDULER') and self.optimizer is not None:
|
||||
self.cfg.LR_SCHEDULER.TOTAL_STEPS = self.max_steps
|
||||
@@ -238,22 +374,22 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
logger=self.logger,
|
||||
optimizer=self.optimizer)
|
||||
|
||||
if self.cfg.DTYPE == 'float16':
|
||||
if self.cfg.DTYPE in ['float16']:
|
||||
if we.is_distributed:
|
||||
if self.use_fairscale:
|
||||
from fairscale.optim.grad_scaler import ShardedGradScaler
|
||||
self.scaler = ShardedGradScaler(enabled=True)
|
||||
elif self.use_fsdp:
|
||||
from torch.distributed.fsdp.sharded_grad_scaler import \
|
||||
ShardedGradScaler
|
||||
self.scaler = ShardedGradScaler()
|
||||
from torch.distributed.fsdp.sharded_grad_scaler import ShardedGradScaler
|
||||
self.scaler = ShardedGradScaler(enabled=True,
|
||||
process_group=None)
|
||||
else:
|
||||
self.scaler = amp.GradScaler()
|
||||
else:
|
||||
self.scaler = amp.GradScaler()
|
||||
else:
|
||||
self.scaler = None
|
||||
|
||||
self.logger.info(self.model)
|
||||
def load_checkpoint(self, checkpoint: dict):
|
||||
"""
|
||||
Load checkpoint function
|
||||
@@ -262,19 +398,56 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
"""
|
||||
if 'model' in checkpoint:
|
||||
if hasattr(self.model, 'module'):
|
||||
self.model.module.load_state_dict(checkpoint['model'])
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model.module, module)
|
||||
if current_module is not None and module in checkpoint[
|
||||
'model']:
|
||||
current_module.load_state_dict(
|
||||
checkpoint['model'][module])
|
||||
self.logger.info(
|
||||
f'Load checkpoint for model.{module}')
|
||||
else:
|
||||
self.model.module.load_state_dict(checkpoint['model'])
|
||||
self.logger.info('Load checkpoint for model.')
|
||||
else:
|
||||
self.model.load_state_dict(checkpoint['model'])
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
self.logger.info(f'Load checkpoint for model.{module}')
|
||||
if current_module is not None and module in checkpoint[
|
||||
'model']:
|
||||
current_module.load_state_dict(
|
||||
checkpoint['model'][module])
|
||||
else:
|
||||
self.model.load_state_dict(checkpoint['model'])
|
||||
self.logger.info('Load checkpoint for model.')
|
||||
else:
|
||||
if hasattr(self.model, 'module'):
|
||||
self.model.module.load_state_dict(checkpoint)
|
||||
else:
|
||||
self.model.load_state_dict(checkpoint)
|
||||
self.logger.info('Load checkpoint for model.')
|
||||
if not self.load_model_only:
|
||||
if 'optimizer' in checkpoint and self.optimizer:
|
||||
self.optimizer.load_state_dict(checkpoint['optimizer'])
|
||||
if self.use_fsdp:
|
||||
for module in self.train_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
if current_module is not None and module in checkpoint[
|
||||
'optimizer']:
|
||||
state = FSDP.optim_state_dict_to_load(
|
||||
current_module, self.optimizer,
|
||||
checkpoint['optimizer'][module])
|
||||
self.optimizer.load_state_dict(state)
|
||||
self.logger.info(
|
||||
f'Load checkpoint for optimizer {module}.')
|
||||
else:
|
||||
self.optimizer.load_state_dict(checkpoint['optimizer'])
|
||||
self.logger.info(f'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 finished.')
|
||||
|
||||
def save_checkpoint(self) -> dict:
|
||||
"""
|
||||
@@ -282,23 +455,76 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
:return:
|
||||
"""
|
||||
ckpt = dict()
|
||||
if not self.use_fsdp and not we.rank == 0:
|
||||
return ckpt
|
||||
ckpt['model'] = OrderedDict()
|
||||
if we.is_distributed:
|
||||
if self.use_fsdp:
|
||||
save_policy = FullStateDictConfig(offload_to_cpu=True,
|
||||
rank0_only=True)
|
||||
with FullyShardedDataParallel.state_dict_type(
|
||||
self.model, StateDictType.FULL_STATE_DICT,
|
||||
save_policy):
|
||||
ckpt['model'] = self.model.state_dict()
|
||||
if self.shard_modules is not None:
|
||||
if self.save_modules is None:
|
||||
self.logger.warning(
|
||||
'When using FSDP, after specifying the modules to be wrapped '
|
||||
'using FSDP_SHARD_MODULES, please set the modules to be saved '
|
||||
'in SAVE_MODULES. If set to None, it means all modules will '
|
||||
'be saved. However, for nested modules, the system cannot '
|
||||
'determine whether they are FSDP instances and need to be '
|
||||
'explicitly set.')
|
||||
assert self.save_modules is not None
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
if current_module is not None:
|
||||
if isinstance(current_module, FSDP):
|
||||
# print(module, current_module._is_root)
|
||||
with FSDP.state_dict_type(
|
||||
current_module,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
save_policy):
|
||||
ckpt['model'][
|
||||
module] = current_module.state_dict()
|
||||
else:
|
||||
ckpt['model'][
|
||||
module] = current_module.state_dict()
|
||||
else:
|
||||
with FSDP.state_dict_type(self.model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
save_policy):
|
||||
ckpt['model'] = self.model.state_dict()
|
||||
else:
|
||||
if hasattr(self.model, 'module'):
|
||||
ckpt['model'] = self.model.module.state_dict()
|
||||
model = self.model.module
|
||||
else:
|
||||
ckpt['model'] = self.model.state_dict()
|
||||
model = self.model
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
if current_module is not None:
|
||||
ckpt['model'][module] = current_module.state_dict()
|
||||
else:
|
||||
ckpt['model'] = model.state_dict()
|
||||
else:
|
||||
ckpt['model'] = self.model.state_dict()
|
||||
if hasattr(self.model, 'module'):
|
||||
model = self.model.module
|
||||
else:
|
||||
model = self.model
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
if current_module is not None:
|
||||
ckpt['model'][module] = current_module.state_dict()
|
||||
else:
|
||||
ckpt['model'] = model.state_dict()
|
||||
if self.optimizer and not self.use_fairscale:
|
||||
ckpt['optimizer'] = self.optimizer.state_dict()
|
||||
if self.use_fsdp and we.is_distributed:
|
||||
ckpt['optimizer'] = OrderedDict()
|
||||
for module in self.train_modules:
|
||||
if hasattr(self.model, module):
|
||||
current_module = getattr(self.model, module)
|
||||
ckpt['optimizer'][module] = FSDP.optim_state_dict(
|
||||
current_module, self.optimizer)
|
||||
else:
|
||||
ckpt['optimizer'] = self.optimizer.state_dict()
|
||||
if self.scaler:
|
||||
ckpt['scaler'] = self.scaler.state_dict()
|
||||
return ckpt
|
||||
@@ -354,7 +580,8 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
})
|
||||
self.current_batch_data[self.mode] = batch_data
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
self.current_batch_data[self.mode].update(
|
||||
self.sample_args.get_lowercase_dict())
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
@@ -395,58 +622,29 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
log_data, log_label, ori_label = [], [], []
|
||||
for result in all_results:
|
||||
# the inference image use
|
||||
ret_images, ret_labels = [], []
|
||||
if 'hint' in result:
|
||||
merge_image = torch.cat([
|
||||
result['hint'][:result['image'].shape[0]], result['image']
|
||||
],
|
||||
dim=2)
|
||||
log_data.append((merge_image.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
else:
|
||||
log_data.append(
|
||||
(result['image'].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append(result['prompt'] + ' NegPrompt: ' +
|
||||
result['n_prompt'])
|
||||
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_labels.append(result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['n_prompt'])
|
||||
log_data.append(ret_images)
|
||||
log_label.append(ret_labels)
|
||||
ori_label.append(result['prompt'])
|
||||
|
||||
self.register_probe({'test_label': log_label})
|
||||
self.register_probe({'eval_label': log_label})
|
||||
self.register_probe({
|
||||
'test_image':
|
||||
'eval_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
|
||||
log_data, log_label, ori_label = [], [], []
|
||||
for result in all_results:
|
||||
# the inference image use
|
||||
if 'train_n_image' in result:
|
||||
if 'hint' in result:
|
||||
merge_image = torch.cat([
|
||||
result['hint'][:result['train_n_image'].shape[0]],
|
||||
result['train_n_image']
|
||||
],
|
||||
dim=2)
|
||||
log_data.append(
|
||||
(merge_image.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
else:
|
||||
log_data.append((result['train_n_image'].permute(
|
||||
1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
|
||||
log_label.append(result['prompt'] + 'NegPrompt' +
|
||||
result['train_n_prompt'])
|
||||
ori_label.append(result['prompt'])
|
||||
if len(log_data) > 0:
|
||||
self.register_probe({'test_train_n_label': log_label})
|
||||
self.register_probe({
|
||||
'test_train_n_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -468,14 +666,23 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
rank=we.rank)
|
||||
all_results.extend(results)
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
log_data, log_label = [], []
|
||||
log_data, log_label, ori_label = [], [], []
|
||||
for result in all_results:
|
||||
# the inference image use
|
||||
log_data.append((result['image'].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append(result['prompt'] +
|
||||
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_labels.append(result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['n_prompt'])
|
||||
log_data.append(ret_images)
|
||||
log_label.append(ret_labels)
|
||||
ori_label.append(result['prompt'])
|
||||
|
||||
self.register_probe({'test_label': log_label})
|
||||
self.register_probe({
|
||||
@@ -485,27 +692,6 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
|
||||
log_data, log_label = [], []
|
||||
for result in all_results:
|
||||
# the inference image use
|
||||
if 'train_n_image' in result:
|
||||
log_data.append(
|
||||
(result['train_n_image'].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append(result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['train_n_prompt'])
|
||||
if len(log_data) > 0:
|
||||
self.register_probe({'test_train_n_label': log_label})
|
||||
self.register_probe({
|
||||
'test_train_n_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
def add_tuner(self, tuner_cfg, model=None):
|
||||
@@ -524,7 +710,7 @@ 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 self.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):
|
||||
@@ -582,6 +768,11 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
freeze_flag = sum([p in name for p in train_part]) > 0
|
||||
if freeze_flag:
|
||||
param.requires_grad = True
|
||||
elif isinstance(train_part, str):
|
||||
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])
|
||||
return model
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -602,7 +793,11 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
return dict_to_yaml('solvername',
|
||||
__class__.__name__,
|
||||
LatentDiffusionSolver.para_dict,
|
||||
set_name=True)
|
||||
set_name=True,
|
||||
exclude_keys=[
|
||||
'EXTRA_KEYS', 'TRAIN_PRECISION', 'MAX_EPOCHS',
|
||||
'NUM_FOLDS'
|
||||
])
|
||||
|
||||
@property
|
||||
def image_out(self):
|
||||
@@ -620,29 +815,32 @@ 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])
|
||||
self.eval_mode()
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
outputs = self.log_image(
|
||||
transfer_data_to_cuda(self.current_batch_data[self.mode]))
|
||||
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)
|
||||
self.train_mode()
|
||||
log_data, log_label = [], []
|
||||
for result in outputs:
|
||||
for result, image in zip(results, images):
|
||||
ret_images, ret_labels = [], []
|
||||
if 'hint' in result:
|
||||
merge_image = torch.cat([
|
||||
result['orig'],
|
||||
result['hint'][:result['orig'].shape[0]],
|
||||
result['recon']
|
||||
],
|
||||
dim=2)
|
||||
else:
|
||||
merge_image = torch.cat([result['orig'], result['recon']],
|
||||
dim=2)
|
||||
log_data.append((merge_image.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append('recon image: ' + result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['n_prompt'])
|
||||
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((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({
|
||||
'train_image':
|
||||
ProbeData(log_data,
|
||||
@@ -651,38 +849,6 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
build_label=log_label)
|
||||
})
|
||||
self.register_probe({'train_label': log_label})
|
||||
|
||||
# the inference image use
|
||||
log_data, log_label = [], []
|
||||
for result in outputs:
|
||||
if 'train_n_image' in result:
|
||||
if 'hint' in result:
|
||||
merge_image = torch.cat([
|
||||
result['orig'],
|
||||
result['hint'][:result['orig'].shape[0]],
|
||||
result['train_n_image']
|
||||
],
|
||||
dim=2)
|
||||
else:
|
||||
merge_image = torch.cat(
|
||||
[result['orig'], result['train_n_image']], dim=2)
|
||||
log_data.append(
|
||||
(merge_image.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append(
|
||||
'recon image: ' + result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['train_n_prompt'])
|
||||
|
||||
if len(log_data) > 0:
|
||||
self.register_probe({'train_n_label': log_label})
|
||||
self.register_probe({
|
||||
'train_n_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
return super().probe_data
|
||||
|
||||
def print_memory_status(self):
|
||||
@@ -719,12 +885,20 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
if logger is None:
|
||||
logger = self.logger
|
||||
train_param_dict = {}
|
||||
forzen_param_dict = {}
|
||||
frozen_param_dict = {}
|
||||
ema_param_dict = {}
|
||||
all_param_numel = 0
|
||||
if we.debug:
|
||||
for key, _ in model.named_modules():
|
||||
logger.info(f'sub modules {key}.')
|
||||
for key, val in model.named_parameters():
|
||||
if 'ema' in key:
|
||||
sub_key = '.'.join(key.split('.', 1)[-1].split('.', 1)[:1])
|
||||
if sub_key in ema_param_dict:
|
||||
ema_param_dict[sub_key] += val.numel()
|
||||
else:
|
||||
ema_param_dict[sub_key] = val.numel()
|
||||
continue
|
||||
if val.requires_grad:
|
||||
sub_key = '.'.join(key.split('.', 1)[-1].split('.', 2)[:2])
|
||||
if sub_key in train_param_dict:
|
||||
@@ -733,20 +907,26 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
train_param_dict[sub_key] = val.numel()
|
||||
else:
|
||||
sub_key = '.'.join(key.split('.', 1)[-1].split('.', 1)[:1])
|
||||
if sub_key in forzen_param_dict:
|
||||
forzen_param_dict[sub_key] += val.numel()
|
||||
if sub_key in frozen_param_dict:
|
||||
frozen_param_dict[sub_key] += val.numel()
|
||||
else:
|
||||
forzen_param_dict[sub_key] = val.numel()
|
||||
frozen_param_dict[sub_key] = val.numel()
|
||||
all_param_numel += val.numel()
|
||||
if we.debug:
|
||||
logger.info(key)
|
||||
train_param_numel = sum(train_param_dict.values())
|
||||
forzen_param_numel = sum(forzen_param_dict.values())
|
||||
frozen_param_numel = sum(frozen_param_dict.values())
|
||||
logger.info(
|
||||
f'Load trainable params {train_param_numel} / {all_param_numel} = '
|
||||
f'{train_param_numel / all_param_numel:.2%}, '
|
||||
f'train part: {train_param_dict}.')
|
||||
logger.info(
|
||||
f'Load forzen params {forzen_param_numel} / {all_param_numel} = '
|
||||
f'{forzen_param_numel / all_param_numel:.2%}, '
|
||||
f'forzen part: {forzen_param_dict}.')
|
||||
f'Load frozen params {frozen_param_numel} / {all_param_numel} = '
|
||||
f'{frozen_param_numel / all_param_numel:.2%}, '
|
||||
f'frozen part: {frozen_param_dict}.')
|
||||
if len(ema_param_dict) > 0:
|
||||
ema_param_numel = sum(ema_param_dict.values())
|
||||
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}.')
|
||||
@@ -1,8 +1,13 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
from scepter.modules.solver.hooks.hook import Hook
|
||||
from scepter.modules.solver.hooks.registry import HOOKS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
@@ -30,6 +35,30 @@ class BackwardHook(Hook):
|
||||
'EMPTY_CACHE_STEP': {
|
||||
'value': -1,
|
||||
'description': 'the memory empty step!'
|
||||
},
|
||||
'DO_PROFILE': {
|
||||
'value': False,
|
||||
'description': 'whether to do profiling!'
|
||||
},
|
||||
'PROFILE_DIR': {
|
||||
'value': None,
|
||||
'description': 'the dir for profiling!'
|
||||
},
|
||||
'PROFILE_WAIT': {
|
||||
'value': 1,
|
||||
'description': 'the wait steps for profiling!'
|
||||
},
|
||||
'PROFILE_WARMUP': {
|
||||
'value': 1,
|
||||
'description': 'the warmup steps for profiling!'
|
||||
},
|
||||
'PROFILE_ACTIVE': {
|
||||
'value': 3,
|
||||
'description': 'the active steps for profiling!'
|
||||
},
|
||||
'REPEAT': {
|
||||
'value': 1,
|
||||
'description': 'the repeat steps for profiling!'
|
||||
}
|
||||
}]
|
||||
|
||||
@@ -40,17 +69,48 @@ class BackwardHook(Hook):
|
||||
self.empty_cache_step = cfg.get('EMPTY_CACHE_STEP', -1)
|
||||
self.accumulate_step = cfg.get('ACCUMULATE_STEP', 1)
|
||||
self.current_step = 0
|
||||
self.wait = cfg.get('PROFILE_WAIT', 1)
|
||||
self.warmup = cfg.get('PROFILE_WARMUP', 1)
|
||||
self.active = cfg.get('PROFILE_ACTIVE', 3)
|
||||
self.repeat = cfg.get('REPEAT', 1)
|
||||
self.do_profile = cfg.get('DO_PROFILE', False)
|
||||
self.profile_dir = cfg.get('PROFILE_DIR', None)
|
||||
self.profile_step = 0
|
||||
self.prof = None
|
||||
|
||||
def before_solve(self, solver):
|
||||
if we.rank != 0:
|
||||
return
|
||||
if self.profile_dir is None:
|
||||
self.log_dir = os.path.join(solver.work_dir, 'profile')
|
||||
self._local_log_dir, _ = FS.map_to_local(self.log_dir)
|
||||
os.makedirs(self._local_log_dir, exist_ok=True)
|
||||
if self.do_profile:
|
||||
self.prof = torch.profiler.profile(
|
||||
schedule=torch.profiler.schedule(wait=self.wait, warmup=self.warmup, active=self.active, repeat=self.repeat),
|
||||
on_trace_ready=torch.profiler.tensorboard_trace_handler(self._local_log_dir),
|
||||
record_shapes=True,
|
||||
with_stack=True)
|
||||
self.prof.start()
|
||||
solver.logger.info(f'Profiler start ...')
|
||||
solver.logger.info(f'Profiler: save to {self.log_dir}')
|
||||
def profile(self, solver):
|
||||
if self.prof is None: return
|
||||
if we.rank == 0 and self.do_profile:
|
||||
if self.profile_step < self.wait + self.warmup + self.active:
|
||||
self.prof.step()
|
||||
self.profile_step += 1
|
||||
else:
|
||||
self.prof.stop()
|
||||
self.do_profile = False
|
||||
solver.logger.info(f'Profiler stop after {self.profile_step} steps')
|
||||
FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir)
|
||||
def grad_clip(self, parameters):
|
||||
torch.nn.utils.clip_grad_norm_(parameters=parameters,
|
||||
max_norm=self.gradient_clip,
|
||||
norm_type=2)
|
||||
|
||||
def after_iter(self, solver):
|
||||
if (hasattr(solver, 'use_fsdp')
|
||||
and solver.use_fsdp) and self.accumulate_step > 1:
|
||||
self.logger.info("Fsdp don't surpport gradient accumulate.")
|
||||
self.accumulate_step = 1
|
||||
if solver.optimizer is not None and solver.is_train_mode:
|
||||
if solver.loss is None:
|
||||
warnings.warn(
|
||||
@@ -58,28 +118,37 @@ class BackwardHook(Hook):
|
||||
)
|
||||
return
|
||||
if solver.scaler is not None:
|
||||
solver.scaler.scale(solver.loss).backward()
|
||||
solver.scaler.scale(solver.loss/self.accumulate_step).backward()
|
||||
if self.gradient_clip > 0:
|
||||
solver.scaler.unscale_(solver.optimizer)
|
||||
self.grad_clip(solver.train_parameters())
|
||||
self.current_step += 1
|
||||
# Suppose profiler run after backward, so we need to set backward_prev_step
|
||||
# as the previous one step before the backward step
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
self.profile(solver)
|
||||
solver.scaler.step(solver.optimizer)
|
||||
solver.scaler.update()
|
||||
solver.optimizer.zero_grad()
|
||||
else:
|
||||
solver.loss.backward()
|
||||
(solver.loss/self.accumulate_step).backward()
|
||||
if self.gradient_clip > 0:
|
||||
self.grad_clip(solver.train_parameters())
|
||||
self.current_step += 1
|
||||
# Suppose profiler run after backward, so we need to set backward_prev_step
|
||||
# as the previous one step before the backward step
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
self.profile(solver)
|
||||
solver.optimizer.step()
|
||||
solver.optimizer.zero_grad()
|
||||
if solver.lr_scheduler:
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
solver.lr_scheduler.step()
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
setattr(solver, 'backward_step', True)
|
||||
self.current_step = 0
|
||||
else:
|
||||
setattr(solver, 'backward_step', False)
|
||||
solver.loss = None
|
||||
if self.empty_cache_step > 0 and solver.total_iter % self.empty_cache_step == 0:
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -13,7 +13,6 @@ from scepter.modules.solver.hooks.registry import HOOKS
|
||||
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 swift import push_to_hub
|
||||
|
||||
_DEFAULT_CHECKPOINT_PRIORITY = 300
|
||||
|
||||
@@ -109,35 +108,42 @@ class CheckpointHook(Hook):
|
||||
or solver.total_iter == solver.max_steps - 1):
|
||||
solver.logger.info(
|
||||
f'Saving checkpoint after {solver.total_iter + 1} steps')
|
||||
if we.rank == 0:
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}.pth'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
if not self.disable_save_snapshot:
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}.pth'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
if not self.disable_save_snapshot:
|
||||
checkpoint = solver.save_checkpoint()
|
||||
if we.rank == 0:
|
||||
with FS.put_to(save_path) as local_path:
|
||||
with open(local_path, 'wb') as f:
|
||||
checkpoint = solver.save_checkpoint()
|
||||
torch.save(checkpoint, f)
|
||||
del checkpoint
|
||||
|
||||
from swift import SwiftModel
|
||||
if isinstance(solver.model, SwiftModel):
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
from swift import SwiftModel
|
||||
if isinstance(solver.model, SwiftModel) or (
|
||||
hasattr(solver.model, 'module')
|
||||
and isinstance(solver.model.module, SwiftModel)):
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
if we.rank == 0:
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
solver.model.save_pretrained(local_folder)
|
||||
if hasattr(solver.model, 'module'):
|
||||
solver.model.module.save_pretrained(local_folder)
|
||||
else:
|
||||
solver.model.save_pretrained(local_folder)
|
||||
FS.put_dir_from_local_dir(local_folder, save_path)
|
||||
else:
|
||||
if hasattr(solver, 'save_pretrained'):
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
FS.make_dir(local_folder)
|
||||
ckpt, cfg = solver.save_pretrained()
|
||||
else:
|
||||
if hasattr(solver, 'save_pretrained'):
|
||||
save_path = osp.join(
|
||||
solver.work_dir, 'checkpoints/{}-{}'.format(
|
||||
self.save_name_prefix, solver.total_iter + 1))
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
FS.make_dir(local_folder)
|
||||
ckpt, cfg = solver.save_pretrained()
|
||||
if we.rank == 0:
|
||||
with FS.put_to(
|
||||
os.path.join(
|
||||
local_folder,
|
||||
@@ -150,6 +156,7 @@ class CheckpointHook(Hook):
|
||||
'configuration.json')) as local_path:
|
||||
json.dump(cfg, open(local_path, 'w'))
|
||||
FS.put_dir_from_local_dir(local_folder, save_path)
|
||||
del ckpt
|
||||
|
||||
if self.save_last and solver.total_iter == solver.max_steps - 1:
|
||||
with FS.get_fs_client(save_path) as client:
|
||||
@@ -219,8 +226,10 @@ class CheckpointHook(Hook):
|
||||
with open(local_file, 'wb') as f:
|
||||
torch.save(checkpoint['pre_state_dict'], f)
|
||||
client.put_object_from_local_file(local_file, save_path)
|
||||
del checkpoint
|
||||
|
||||
def after_all_iter(self, solver):
|
||||
from swift import push_to_hub
|
||||
if we.rank == 0:
|
||||
if self.push_to_hub and self.last_ckpt:
|
||||
with FS.get_dir_to_local_dir(self.last_ckpt) as local_dir:
|
||||
|
||||
@@ -23,6 +23,26 @@ class ProbeDataHook(Hook):
|
||||
'PROB_INTERVAL': {
|
||||
'value': 1000,
|
||||
'description': 'the interval for log print!'
|
||||
},
|
||||
'SAVE_NAME_PREFIX': {
|
||||
'value': 'step',
|
||||
'description': 'the prefix for save name!'
|
||||
},
|
||||
'SAVE_PROBE_PREFIX': {
|
||||
'value': None,
|
||||
'description': 'the prefix for save probe!'
|
||||
},
|
||||
'SAVE_LAST': {
|
||||
'value': False,
|
||||
'description': 'whether to save last!'
|
||||
},
|
||||
'SAVE_IMAGE_POSTFIX': {
|
||||
'value': 'jpg',
|
||||
'description': 'the postfix for save image!'
|
||||
},
|
||||
'SAVE_VIDEO_POSTFIX': {
|
||||
'value': 'mp4',
|
||||
'description': 'the postfix for save video!'
|
||||
}
|
||||
}]
|
||||
|
||||
@@ -34,31 +54,46 @@ class ProbeDataHook(Hook):
|
||||
self.save_probe_prefix = cfg.get('SAVE_PROBE_PREFIX', None)
|
||||
self.save_last = cfg.get('SAVE_LAST', False)
|
||||
self.save_image_postfix = cfg.get('SAVE_IMAGE_POSTFIX', 'jpg')
|
||||
self.save_video_postfix = cfg.get('SAVE_VIDEO_POSTFIX', 'mp4')
|
||||
|
||||
def before_all_iter(self, solver):
|
||||
pass
|
||||
|
||||
if not solver.mode == 'train' and hasattr(solver, 'eval_interval'):
|
||||
solver.eval_interval = self.prob_interval
|
||||
def before_iter(self, solver):
|
||||
pass
|
||||
|
||||
def get_key_level_prefix(self, key, save_folder, total_iter):
|
||||
if self.save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder,
|
||||
self.save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(
|
||||
save_folder,
|
||||
key.replace('/', '_') + f'_step_{total_iter}')
|
||||
return ret_prefix
|
||||
def after_iter(self, solver):
|
||||
if solver.mode == 'train' and solver.total_iter % self.prob_interval == 0:
|
||||
probe_dict = solver.probe_data
|
||||
if we.rank == 0:
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{solver.total_iter}'
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{solver.total_iter}'
|
||||
)
|
||||
setattr(solver,
|
||||
f'{solver.mode}_pre_save_paras',
|
||||
{
|
||||
"save_folder":save_folder,
|
||||
"save_probe_prefix": self.save_probe_prefix,
|
||||
"image_postfix": self.save_image_postfix,
|
||||
"video_postfix": self.save_video_postfix,
|
||||
"step": solver.total_iter,
|
||||
}
|
||||
)
|
||||
gather_probe_dict = solver.probe_data
|
||||
if we.rank == 0:
|
||||
ret_data = {}
|
||||
for k, v in probe_dict.items():
|
||||
if self.save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder,
|
||||
self.save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(
|
||||
save_folder,
|
||||
k.replace('/', '_') + f'_step_{solver.total_iter}')
|
||||
ret_one = v.to_log(ret_prefix, self.save_image_postfix)
|
||||
for k, v in gather_probe_dict.items():
|
||||
ret_prefix = self.get_key_level_prefix(k, save_folder, solver.total_iter)
|
||||
ret_one = v.to_log(ret_prefix, image_postfix=self.save_image_postfix,
|
||||
video_postfix=self.save_video_postfix)
|
||||
if (isinstance(ret_one, list)
|
||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||
continue
|
||||
@@ -74,23 +109,29 @@ class ProbeDataHook(Hook):
|
||||
|
||||
def after_all_iter(self, solver):
|
||||
if not solver.mode == 'train':
|
||||
step = solver._total_iter[
|
||||
'train'] if 'train' in solver._total_iter else 0
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{step}'
|
||||
)
|
||||
setattr(solver,
|
||||
f'{solver.mode}_pre_save_paras',
|
||||
{
|
||||
"save_folder": save_folder,
|
||||
"save_probe_prefix": self.save_probe_prefix,
|
||||
"image_postfix": self.save_image_postfix,
|
||||
"step": step,
|
||||
"video_postfix": self.save_video_postfix
|
||||
}
|
||||
)
|
||||
probe_dict = solver.probe_data
|
||||
if we.rank == 0:
|
||||
step = solver._total_iter[
|
||||
'train'] if 'train' in solver._total_iter else 0
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{step}')
|
||||
ret_data = {}
|
||||
for k, v in probe_dict.items():
|
||||
if self.save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder,
|
||||
self.save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(
|
||||
save_folder,
|
||||
k.replace('/', '_') + f'_step_{step}')
|
||||
ret_one = v.to_log(ret_prefix, self.save_image_postfix)
|
||||
ret_prefix = self.get_key_level_prefix(k, save_folder, step)
|
||||
ret_one = v.to_log(ret_prefix, image_postfix=self.save_image_postfix,
|
||||
video_postfix=self.save_video_postfix)
|
||||
if (isinstance(ret_one, list)
|
||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||
continue
|
||||
|
||||
@@ -124,12 +124,17 @@ class LogHook(Hook):
|
||||
|
||||
self.time = time.time()
|
||||
self.start_time = time.time()
|
||||
self.all_throughput = 0
|
||||
self.data_time = 0
|
||||
self.batch_size = defaultdict(dict)
|
||||
|
||||
def before_all_iter(self, solver):
|
||||
self.time = time.time()
|
||||
self.last_log_step = (solver.mode, 0)
|
||||
|
||||
if hasattr(solver, "datas"):
|
||||
for k, v in solver.datas.items():
|
||||
if hasattr(v, 'batch_size'):
|
||||
self.batch_size[k] = v.batch_size
|
||||
def before_iter(self, solver):
|
||||
data_time = time.time() - self.time
|
||||
self.data_time = data_time
|
||||
@@ -141,16 +146,21 @@ class LogHook(Hook):
|
||||
outputs = solver.iter_outputs.copy()
|
||||
outputs['time'] = iter_time
|
||||
outputs['data_time'] = self.data_time
|
||||
if 'batch_size' in outputs:
|
||||
batch_size = outputs.pop('batch_size')
|
||||
else:
|
||||
batch_size = 1
|
||||
if solver.mode in self.batch_size:
|
||||
outputs['throughput'] = int(self.batch_size[solver.mode] * we.world_size / iter_time * 86400)
|
||||
log_agg.update(outputs, 1)
|
||||
log_agg = log_agg.aggregate(self.log_interval)
|
||||
if 'throughput' in log_agg:
|
||||
log_agg['throughput'] = f"{int(log_agg['throughput'][-1])}/day"
|
||||
if solver.mode in self.batch_size:
|
||||
log_agg['all_throughput'] = (solver.iter + 1) * we.world_size * self.batch_size[solver.mode]
|
||||
|
||||
if self.show_gpu_mem:
|
||||
outputs['nvidia-smi'] = print_memory_status()
|
||||
log_agg.update(outputs, batch_size)
|
||||
log_agg['nvidia-smi'] = str(print_memory_status()) +"MiB"
|
||||
|
||||
if (solver.iter + 1) % self.log_interval == 0:
|
||||
_print_iter_log(solver,
|
||||
log_agg.aggregate(self.log_interval),
|
||||
log_agg,
|
||||
start_time=self.start_time,
|
||||
mode=solver.mode)
|
||||
self.last_log_step = (solver.mode, solver.iter + 1)
|
||||
|
||||
Reference in New Issue
Block a user