File size: 4,542 Bytes
4e2a1b3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 | import os, torch
from accelerate import Accelerator
class TensorBoardLogger:
def __init__(self, log_dir):
from torch.utils.tensorboard import SummaryWriter
self.writer = SummaryWriter(log_dir=log_dir)
print(f"TensorBoard is enabled. Run `tensorboard --logdir={log_dir}` to visualize the training progress.")
def log(self, key, value, step):
self.writer.add_scalar(key, value, step)
def close(self):
if self.writer is not None:
self.writer.close()
class SwanLabLogger:
def __init__(self, project_name="DiffSynth-Studio", log_dir=None):
import swanlab
project_name = os.environ.get("SWANLAB_PROJECT", project_name)
self.swanlab = swanlab
self.swanlab.init(project=project_name, logdir=log_dir)
print(f"SwanLab is enabled. Project: {project_name}")
def log(self, key, value, step):
self.swanlab.log({key: value}, step=step)
def close(self):
self.swanlab.finish()
class WandbLogger:
def __init__(self, project_name="DiffSynth-Studio", log_dir=None):
import wandb
project_name = os.environ.get("WANDB_PROJECT", project_name)
self.wandb = wandb
self.run = self.wandb.init(project=project_name, dir=log_dir)
print(f"Wandb is enabled. Project: {project_name}")
def log(self, key, value, step):
self.wandb.log({key: value}, step=step)
def close(self):
self.wandb.finish()
class ModelLogger:
def __init__(
self, output_path, remove_prefix_in_ckpt=None, state_dict_converter=lambda x: x,
enable_tensorboard_log=False,
enable_swanlab_log=False, swanlab_project="DiffSynth-Studio",
enable_wandb_log=False, wandb_project="DiffSynth-Studio",
):
self.output_path = output_path
self.remove_prefix_in_ckpt = remove_prefix_in_ckpt
self.state_dict_converter = state_dict_converter
self.num_steps = 0
# Loggers
self.enable_tensorboard_log = enable_tensorboard_log
self.enable_swanlab_log = enable_swanlab_log
self.swanlab_project = swanlab_project
self.enable_wandb_log = enable_wandb_log
self.wandb_project = wandb_project
self.loggers = []
self.loggers_initialized = False
def init_loggers(self):
if self.enable_tensorboard_log:
self.loggers.append(TensorBoardLogger(os.path.join(self.output_path, "tensorboard_log")))
if self.enable_swanlab_log:
self.loggers.append(SwanLabLogger(project_name=self.swanlab_project, log_dir=os.path.join(self.output_path, "swanlab_log")))
if self.enable_wandb_log:
self.loggers.append(WandbLogger(project_name=self.wandb_project, log_dir=os.path.join(self.output_path, "wandb_log")))
self.loggers_initialized = True
def on_step_end(self, accelerator: Accelerator, model: torch.nn.Module, save_steps=None, **kwargs):
self.num_steps += 1
if accelerator.is_main_process:
if not self.loggers_initialized:
self.init_loggers()
loss = kwargs.get("loss")
if loss is not None:
for logger in self.loggers:
logger.log("loss", loss, self.num_steps)
if save_steps is not None and self.num_steps % save_steps == 0:
self.save_model(accelerator, model, f"step-{self.num_steps}.safetensors")
def on_epoch_end(self, accelerator: Accelerator, model: torch.nn.Module, epoch_id):
self.save_model(accelerator, model, f"epoch-{epoch_id}.safetensors")
def on_training_end(self, accelerator: Accelerator, model: torch.nn.Module, save_steps=None):
if save_steps is not None and self.num_steps % save_steps != 0:
self.save_model(accelerator, model, f"step-{self.num_steps}.safetensors")
for logger in self.loggers:
logger.close()
def save_model(self, accelerator: Accelerator, model: torch.nn.Module, file_name):
accelerator.wait_for_everyone()
state_dict = accelerator.get_state_dict(model)
if accelerator.is_main_process:
state_dict = accelerator.unwrap_model(model).export_trainable_state_dict(state_dict, remove_prefix=self.remove_prefix_in_ckpt)
state_dict = self.state_dict_converter(state_dict)
os.makedirs(self.output_path, exist_ok=True)
path = os.path.join(self.output_path, file_name)
accelerator.save(state_dict, path, safe_serialization=True)
|