chore: import upstream snapshot with attribution
Lint test / lint (push) Has been cancelled

This commit is contained in:
wehub-resource-sync
2026-07-13 13:34:58 +08:00
commit a203934033
1368 changed files with 175001 additions and 0 deletions
+3
View File
@@ -0,0 +1,3 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
from .base import MegatronCallback
from .mapping import megatron_callbacks_map
+46
View File
@@ -0,0 +1,46 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from swift.megatron.trainers import BaseMegatronTrainer
class MegatronCallback:
def __init__(self, trainer: 'BaseMegatronTrainer'):
self.trainer = trainer
self.args = trainer.args
self.state = trainer.state
def on_train_begin(self):
pass
def on_train_end(self):
pass
def on_step_begin(self):
pass
def on_step_end(self):
pass
def on_log(self, logs):
pass
def on_eval_begin(self):
pass
def on_eval_end(self):
pass
def on_eval_step(self):
pass
def on_save(self, output_dir):
"""Called after save_checkpoint() returns.
Note: When async_save is enabled, the checkpoint may not be fully
written to disk yet. Use only for non-I/O-dependent logic, or
ensure async_save is disabled if you need to read the checkpoint.
"""
pass
+43
View File
@@ -0,0 +1,43 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
import gc
from .base import MegatronCallback
class DefaultFlowCallback(MegatronCallback):
def on_train_begin(self):
args = self.args
if args.manual_gc:
gc.disable()
gc.collect()
def on_step_end(self):
args = self.args
state = self.state
state.consumed_train_samples += args.global_batch_size
if state.iteration == 1 or state.iteration % args.logging_steps == 0:
state.should_log = True
if args.eval_steps and state.iteration % args.eval_steps == 0 and args.eval_iters > 0:
state.should_eval = True
if args.save_steps and state.iteration % args.save_steps == 0:
state.should_save = True
if state.iteration >= args.train_iters:
if args.eval_iters > 0:
state.should_eval = True
state.should_save = True
if args.manual_gc and args.manual_gc_steps != 0 and state.iteration % args.manual_gc_steps == 0:
gc.collect()
def on_eval_begin(self):
args = self.args
if args.manual_gc and args.manual_gc_eval:
gc.collect()
def on_eval_end(self):
args = self.args
if args.manual_gc and args.manual_gc_eval:
gc.collect(generation=0)
+14
View File
@@ -0,0 +1,14 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
from .default_flow import DefaultFlowCallback
from .print import PrintCallback
from .swanlab import SwanlabCallback
from .tensorboard import TensorboardCallback
from .wandb import WandbCallback
megatron_callbacks_map = {
'print': PrintCallback,
'default_flow': DefaultFlowCallback,
'swanlab': SwanlabCallback,
'wandb': WandbCallback,
'tensorboard': TensorboardCallback,
}
+69
View File
@@ -0,0 +1,69 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
import os
import time
import torch
from tqdm import tqdm
from swift.megatron.utils import reduce_max_stat_across_model_parallel_group
from swift.utils import JsonlWriter, format_time, get_logger, is_last_rank
from .base import MegatronCallback
logger = get_logger()
class PrintCallback(MegatronCallback):
def __init__(self, trainer):
super().__init__(trainer)
self.training_bar = None
self.eval_bar = None
self.jsonl_writer = None
self.is_write_rank = is_last_rank()
def on_train_begin(self):
self.training_bar = tqdm(
total=self.args.train_iters, dynamic_ncols=True, disable=not self.is_write_rank, desc='Train: ')
self.start_step = self.state.iteration
self.training_bar.update(self.state.iteration)
self.current_step = self.state.iteration
self.start_time = time.time()
logging_path = os.path.join(self.args.output_dir, 'logging.jsonl')
logger.info(f'logging_path: {logging_path}')
self.jsonl_writer = JsonlWriter(logging_path, enable_async=True, write_on_rank='last')
def on_train_end(self):
self.training_bar.close()
self.training_bar = None
def on_step_end(self):
n_step = self.state.iteration - self.current_step
self.current_step = self.state.iteration
self.training_bar.update(n_step)
def on_eval_begin(self):
self.eval_bar = tqdm(
total=self.args.eval_iters, dynamic_ncols=True, disable=not self.is_write_rank, desc='Evaluate: ')
def on_eval_end(self):
self.eval_bar.close()
self.eval_bar = None
def on_eval_step(self):
self.eval_bar.update()
def on_log(self, logs):
state = self.state
args = self.args
logs['iteration'] = f'{state.iteration}/{args.train_iters}'
elapsed = time.time() - self.start_time
logs['elapsed_time'] = format_time(elapsed)
n_steps = state.iteration - self.start_step
train_speed = elapsed / n_steps if n_steps > 0 else 0.0
logs['remaining_time'] = format_time((args.train_iters - state.iteration) * train_speed)
memory = reduce_max_stat_across_model_parallel_group(torch.cuda.max_memory_reserved() / 1024**3)
logs['memory(GiB)'] = round(memory, 2)
logs['train_speed(s/it)'] = round(train_speed, 6)
logs = {k: round(v, 8) if isinstance(v, float) else v for k, v in logs.items()}
self.jsonl_writer.append(logs)
if self.is_write_rank:
self.training_bar.write(str(logs))
+35
View File
@@ -0,0 +1,35 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
import os
from swift.utils import check_json_format, is_last_rank
from .base import MegatronCallback
from .utils import rewrite_logs
class SwanlabCallback(MegatronCallback):
def __init__(self, trainer):
super().__init__(trainer)
args = self.args
self.config = check_json_format(vars(args))
if args.swanlab_exp_name is None:
args.swanlab_exp_name = args.output_dir
self.save_dir = os.path.join(args.output_dir, 'swanlab')
self.writer = None
self.setup()
def setup(self):
import swanlab
args = self.args
if is_last_rank():
swanlab.init(
logdir=self.save_dir,
experiment_name=args.swanlab_exp_name,
project=args.swanlab_project,
config=self.config)
self.writer = swanlab
def on_log(self, logs):
logs = rewrite_logs(logs)
if is_last_rank():
self.writer.log(logs, step=self.state.iteration)
+32
View File
@@ -0,0 +1,32 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
from swift.utils import check_json_format, is_last_rank
from .base import MegatronCallback
from .utils import rewrite_logs
class TensorboardCallback(MegatronCallback):
def __init__(self, trainer):
super().__init__(trainer)
args = self.args
self.config = check_json_format(vars(args))
self.save_dir = args.tensorboard_dir
if self.save_dir is None:
self.save_dir = f'{args.output_dir}/runs'
from torch.utils.tensorboard import SummaryWriter
self.writer = None
if is_last_rank():
self.writer = SummaryWriter(log_dir=self.save_dir, max_queue=args.tensorboard_queue_size)
for k, v in self.config.items():
self.writer.add_text(k, str(v), global_step=self.state.iteration)
def on_log(self, logs):
logs = rewrite_logs(logs)
if self.writer:
for k, v in logs.items():
self.writer.add_scalar(k, v, self.state.iteration)
def on_train_end(self):
if self.writer:
self.writer.close()
self.writer = None
+19
View File
@@ -0,0 +1,19 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
def rewrite_logs(logs):
new_logs = {}
for k, v in logs.items():
if isinstance(v, str):
continue
k = k.replace('/', '_')
if k.startswith('eval_'):
k = k[len('eval_'):]
k = f'eval/{k}'
elif k.startswith('test_'):
k = k[len('test_'):]
k = f'test/{k}'
else:
k = f'train/{k}'
new_logs[k] = v
return new_logs
+31
View File
@@ -0,0 +1,31 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
import os
from swift.utils import check_json_format, is_last_rank
from .base import MegatronCallback
from .utils import rewrite_logs
class WandbCallback(MegatronCallback):
def __init__(self, trainer):
super().__init__(trainer)
args = self.args
self.config = check_json_format(vars(args))
if args.wandb_exp_name is None:
args.wandb_exp_name = args.output_dir
self.save_dir = args.output_dir
self.writer = None
self.setup()
def setup(self):
import wandb
args = self.args
if is_last_rank():
wandb.init(dir=self.save_dir, name=args.wandb_exp_name, project=args.wandb_project, config=self.config)
self.writer = wandb
def on_log(self, logs):
logs = rewrite_logs(logs)
if is_last_rank():
self.writer.log(logs, step=self.state.iteration)