35 lines
730 B
Python
35 lines
730 B
Python
"""Distributed trainer package exports."""
|
|||
|
|
|
||
|
|
from wall_x.trainer.fsdp_trainer.base_trainer import (
|
||
|
|
BaseDistributedTrainer,
|
||
|
|
DistributedTrainer,
|
||
|
|
all_gather,
|
||
|
|
all_reduce,
|
||
|
|
barrier,
|
||
|
|
cleanup_distributed,
|
||
|
|
get_local_rank,
|
||
|
|
get_rank,
|
||
|
|
get_world_size,
|
||
|
|
is_main_process,
|
||
|
|
setup_distributed,
|
||
|
|
)
|
||
|
|
from wall_x.trainer.fsdp_trainer.fsdp_trainer import FSDPTrainer
|
||
|
|
|
||
|
|
__all__ = [
|
||
|
|
# Base classes and utilities
|
||
|
|
"BaseDistributedTrainer",
|
||
|
|
"DistributedTrainer",
|
||
|
|
"setup_distributed",
|
||
|
|
"cleanup_distributed",
|
||
|
|
"get_rank",
|
||
|
|
"get_world_size",
|
||
|
|
"get_local_rank",
|
||
|
|
"is_main_process",
|
||
|
|
"barrier",
|
||
|
|
"all_reduce",
|
||
|
|
"all_gather",
|
||
|
|
# Trainers
|
||
|
|
# "DDPTrainer",
|
||
|
|
"FSDPTrainer",
|
||
|
|
]
|