update v0.0.4
This commit is contained in:
@@ -333,12 +333,12 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
data_iter = iter(self.datas[self._mode].dataloader)
|
||||
self.print_memory_status()
|
||||
for step in range(self.max_steps):
|
||||
if 'eval' in self._mode_set and (step % self.eval_interval == 0
|
||||
or step == self.max_steps - 1):
|
||||
if 'eval' in self._mode_set and (self.eval_interval > 0 and
|
||||
step % self.eval_interval == 0):
|
||||
self.run_eval()
|
||||
self.train_mode()
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
batch_data = next(data_iter)
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
if 'meta' in batch_data:
|
||||
@@ -364,6 +364,9 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
if we.debug:
|
||||
self.print_trainable_params_status(prefix='model.')
|
||||
if 'eval' in self._mode_set and (self.eval_interval > 0
|
||||
and step == self.max_steps - 1):
|
||||
self.run_eval()
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@torch.no_grad()
|
||||
|
||||
@@ -4,12 +4,14 @@
|
||||
from scepter.modules.solver.hooks.backward import BackwardHook
|
||||
from scepter.modules.solver.hooks.checkpoint import CheckpointHook
|
||||
from scepter.modules.solver.hooks.data_probe import ProbeDataHook
|
||||
from scepter.modules.solver.hooks.ema import ModelEmaHook
|
||||
from scepter.modules.solver.hooks.hook import Hook
|
||||
from scepter.modules.solver.hooks.log import LogHook, TensorboardLogHook
|
||||
from scepter.modules.solver.hooks.lr import LrHook
|
||||
from scepter.modules.solver.hooks.registry import HOOKS
|
||||
from scepter.modules.solver.hooks.safetensors import SafetensorsHook
|
||||
from scepter.modules.solver.hooks.sampler import DistSamplerHook
|
||||
|
||||
"""
|
||||
Normally, hooks have priorities, below we recommend priority that runs fine (low score MEANS high priority)
|
||||
BackwardHook: 0
|
||||
@@ -47,5 +49,6 @@ after solve:
|
||||
|
||||
__all__ = [
|
||||
'HOOKS', 'BackwardHook', 'CheckpointHook', 'Hook', 'LrHook', 'LogHook',
|
||||
'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook', 'SafetensorsHook'
|
||||
'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook',
|
||||
'SafetensorsHook', 'ModelEmaHook'
|
||||
]
|
||||
|
||||
@@ -147,7 +147,9 @@ class CheckpointHook(Hook):
|
||||
|
||||
if self.save_last and solver.total_iter == solver.max_steps - 1:
|
||||
with FS.get_fs_client(save_path) as client:
|
||||
last_path = osp.join(solver.work_dir, 'checkpoint.pth')
|
||||
last_path = osp.join(
|
||||
solver.work_dir,
|
||||
f'checkpoints/{self.save_name_prefix}-last')
|
||||
client.make_link(last_path, save_path)
|
||||
self.last_ckpt = save_path
|
||||
|
||||
|
||||
@@ -30,7 +30,10 @@ class ProbeDataHook(Hook):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(ProbeDataHook, self).__init__(cfg, logger=logger)
|
||||
self.priority = cfg.get('PRIORITY', _DEFAULT_PROBE_PRIORITY)
|
||||
self.log_interval = cfg.get('PROB_INTERVAL', 1000)
|
||||
self.prob_interval = cfg.get('PROB_INTERVAL', 1000)
|
||||
self.save_name_prefix = cfg.get('SAVE_NAME_PREFIX', 'step')
|
||||
self.save_probe_prefix = cfg.get('SAVE_PROBE_PREFIX', None)
|
||||
self.save_last = cfg.get('SAVE_LAST', False)
|
||||
|
||||
def before_all_iter(self, solver):
|
||||
pass
|
||||
@@ -39,19 +42,23 @@ class ProbeDataHook(Hook):
|
||||
pass
|
||||
|
||||
def after_iter(self, solver):
|
||||
if solver.mode == 'train' and solver.total_iter % self.log_interval == 0:
|
||||
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/step_{solver.total_iter}')
|
||||
f'{solver.mode}_probe/{self.save_probe_prefix}-{solver.total_iter}'
|
||||
)
|
||||
ret_data = {}
|
||||
for k, v in probe_dict.items():
|
||||
ret_one = v.to_log(
|
||||
os.path.join(
|
||||
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}'))
|
||||
k.replace('/', '_') + f'_step_{solver.total_iter}')
|
||||
ret_one = v.to_log(ret_prefix)
|
||||
if (isinstance(ret_one, list)
|
||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||
continue
|
||||
@@ -71,13 +78,19 @@ class ProbeDataHook(Hook):
|
||||
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/step_{step}')
|
||||
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():
|
||||
ret_one = v.to_log(
|
||||
os.path.join(save_folder,
|
||||
k.replace('/', '_') + f'_step_{step}'))
|
||||
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)
|
||||
if (isinstance(ret_one, list)
|
||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||
continue
|
||||
@@ -87,6 +100,16 @@ class ProbeDataHook(Hook):
|
||||
json.dump(ret_data,
|
||||
open(local_path, 'w'),
|
||||
ensure_ascii=False)
|
||||
|
||||
if self.save_last and step == solver.max_steps:
|
||||
with FS.get_fs_client(save_folder) as client:
|
||||
last_save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-last'
|
||||
)
|
||||
print(last_save_folder, save_folder)
|
||||
client.make_link(last_save_folder, save_folder)
|
||||
|
||||
solver.clear_probe()
|
||||
torch.cuda.synchronize()
|
||||
barrier()
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
import torch
|
||||
from torch.distributed.fsdp import (FullStateDictConfig,
|
||||
FullyShardedDataParallel, StateDictType)
|
||||
|
||||
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
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
|
||||
@HOOKS.register_class()
|
||||
class ModelEmaHook(Hook):
|
||||
para_dict = [{
|
||||
'PRIORITY': {
|
||||
'value': 100,
|
||||
'description': 'the priority for processing!'
|
||||
},
|
||||
'BETA': {
|
||||
'value': 0.9999,
|
||||
'description': ''
|
||||
},
|
||||
}]
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(ModelEmaHook, self).__init__(cfg, logger=logger)
|
||||
self.priority = cfg.get('PRIORITY', 100)
|
||||
self.beta = cfg.get('BETA', 0.9999)
|
||||
|
||||
def after_iter(self, solver):
|
||||
if solver.model.use_ema:
|
||||
model_ema = solver.model.model_ema
|
||||
model = solver.model.model
|
||||
self.ema(model_ema, model, use_fsdp=solver.use_fsdp)
|
||||
|
||||
@torch.no_grad()
|
||||
def ema(self, net_ema, net, use_fsdp=True):
|
||||
if we.is_distributed:
|
||||
if use_fsdp:
|
||||
save_policy = FullStateDictConfig(offload_to_cpu=False,
|
||||
rank0_only=False)
|
||||
with FullyShardedDataParallel.state_dict_type(
|
||||
net, StateDictType.FULL_STATE_DICT, save_policy):
|
||||
nonema_state = net.state_dict()
|
||||
elif hasattr(net, 'module'):
|
||||
nonema_state = net.module.state_dict()
|
||||
else:
|
||||
nonema_state = net.state_dict()
|
||||
else:
|
||||
nonema_state = net.state_dict()
|
||||
|
||||
for k, v in net_ema.named_parameters():
|
||||
v.copy_(nonema_state[k].lerp(v, self.beta))
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('hook',
|
||||
__class__.__name__,
|
||||
ModelEmaHook.para_dict,
|
||||
set_name=True)
|
||||
Reference in New Issue
Block a user