-
Notifications
You must be signed in to change notification settings - Fork 9
Expand file tree
/
Copy pathmain_diff_rloo_trainer.py
More file actions
77 lines (61 loc) · 2.67 KB
/
Copy pathmain_diff_rloo_trainer.py
File metadata and controls
77 lines (61 loc) · 2.67 KB
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
import logging
import os
import pathlib
import hydra
import torch
import transformers
from omegaconf import OmegaConf
# WARNING: the args for trainer are from the official Arguments, do not refer to that in our config.py
from src.train.callbacks import DiffusionWandbCallback
from src.train.config import ConfigPathArguments, CustomRLOOConfig
from src.train.rloo_trainer import CommonRLOOTrainer
from src.train.train_utilis import setup_debug
# rank = os.environ.get("RANK")
# setup_debug(int(rank))
log_format = "%(asctime)s - %(name)s - %(levelname)s - %(message)s"
logging.basicConfig(level=logging.INFO, format=log_format)
logger = logging.getLogger(__name__)
def train(cfg, training_args):
model = hydra.utils.instantiate(
OmegaConf.load(cfg.model_config),
init_alpha=training_args.init_alpha,
init_beta=training_args.init_beta,
relative=training_args.relative,
prediction_type=training_args.prediction_type,
fsdp=training_args.fsdp,
max_inference_steps=training_args.max_inference_steps,
)
logger.info(f"model loaded from {cfg.model_config}")
reward_model = hydra.utils.instantiate(OmegaConf.load(cfg.reward_model_config)).eval()
logger.info(f"reward model loaded from {cfg.reward_model_config}")
train_dataset = hydra.utils.instantiate(OmegaConf.load(cfg.train_dataset))
logger.info(f"train dataset loaded from {cfg.train_dataset}")
data_collator = hydra.utils.instantiate(OmegaConf.load(cfg.data_collator))
logger.info(f"data collator loaded from {cfg.data_collator}")
# TODO: someone may prefer save training_args in a file
trainer = CommonRLOOTrainer(
config=training_args,
policy=model,
reward_model=reward_model,
data_collator=data_collator,
train_dataset=train_dataset,
eval_dataset=train_dataset,
)
if "wandb" in training_args.report_to:
wandb_callback = DiffusionWandbCallback(trainer=trainer)
logger.info("wandb callback added")
trainer.add_callback(wandb_callback)
if (
list(pathlib.Path(training_args.output_dir).glob("checkpoint-*")) and args.resume_from_checkpoint is not None
) or os.path.isdir(args.resume_from_checkpoint):
# jugde whether resume_from_checkpoint is a path
if os.path.isdir(args.resume_from_checkpoint):
trainer.train(resume_from_checkpoint=args.resume_from_checkpoint)
else:
trainer.train(resume_from_checkpoint=True)
else:
trainer.train()
if __name__ == "__main__":
parser = transformers.HfArgumentParser((ConfigPathArguments, CustomRLOOConfig))
cfg, args = parser.parse_args_into_dataclasses()
train(cfg, args)