update v0.0.4

This commit is contained in:
LouieStark
2024-03-31 13:08:41 +08:00
parent 35aada8ce8
commit bf53829530
106 changed files with 6927 additions and 889 deletions
+6 -3
View File
@@ -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 -1
View File
@@ -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'
]
+3 -1
View File
@@ -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
+35 -12
View File
@@ -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()
+60
View File
@@ -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)