-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
189 lines (158 loc) · 6.85 KB
/
Copy pathtrain.py
File metadata and controls
189 lines (158 loc) · 6.85 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
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
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
import math
import argparse
from pathlib import Path
from loguru import logger
import time
import os
import warnings
import torch
import pytorch_lightning as pl
from pytorch_lightning.loggers import TensorBoardLogger
from pytorch_lightning.callbacks import ModelCheckpoint, LearningRateMonitor
# 过滤 Albumentations 的 blur_limit 警告
warnings.filterwarnings("ignore", category=UserWarning, module="albumentations")
from src.config.default import get_cfg_defaults
from src.utils.misc import setup_gpus
from src.lightning.data import MultiSceneDataModule
from src.lightning.lightning_crft import PL_CRFT
# 开启 TF32 加速矩阵运算
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
def parse_gpu_args(gpu_arg):
"""
解析 GPU 参数,支持:
- int 类型: 1 (代表使用 1 张卡,默认为 0号卡)
- str 类型: "0" / "1" (指定某张卡) 或 "0,1" (使用双卡 DDP)
"""
if isinstance(gpu_arg, int):
devices = list(range(gpu_arg))
elif isinstance(gpu_arg, str):
if ',' in gpu_arg:
devices = [int(x.strip()) for x in gpu_arg.split(',') if x.strip() != '']
else:
devices = [int(gpu_arg.strip())]
else:
devices = [0]
return devices
def main():
parser = argparse.ArgumentParser(description="CRFT Training Script with Multi-GPU Support")
parser.add_argument('--exp_name', type=str, default='crft-train', help='实验名称 (建议不同对比实验用不同名字)')
parser.add_argument('--batch_size', type=int, default=16, help='单卡 Batch Size')
parser.add_argument('--epochs', type=int, default=30, help='训练轮数 Epochs')
parser.add_argument('--gpus', type=str, default='0', help='指定 GPU,例如 "0"、"1"(单卡独立跑)或 "0,1"(双卡分布式)')
parser.add_argument('--num_workers', type=int, default=8, help='DataLoader 线程数')
parser.add_argument('--pretrained_ckpt', type=str, default='./pretrain/CRFT_OSdataset.ckpt', help='预训练权重路径')
args = parser.parse_args()
# 解析请求的 GPU
devices = parse_gpu_args(args.gpus)
num_gpus = len(devices)
print("=" * 60)
print("Starting Training")
print(f" - Experiment Name: {args.exp_name}")
print(f" - Single GPU Batch size: {args.batch_size}")
print(f" - Epochs: {args.epochs}")
print(f" - Target GPUs: {devices} (Total: {num_gpus} GPU(s))")
print(f" - Pretrained checkpoint: {args.pretrained_ckpt}")
print("=" * 60)
# 加载配置文件
config = get_cfg_defaults()
config.merge_from_file('configs/crft/outdoor/visible_thermal.py')
config.merge_from_file('configs/data/osdataset_640.py')
# 覆盖配置选项
config.TRAINER.USE_WANDB = False
config.TRAINER.MAX_EPOCHS = args.epochs
pl.seed_everything(config.TRAINER.SEED, workers=True)
# 🌟 修改点 2:配置分布式策略 (Strategy) 与设备分配
if torch.cuda.is_available() and num_gpus > 0:
accelerator = 'gpu'
if num_gpus > 1:
# 双卡训练时开启 DDP 分布式并行策略
strategy = 'ddp_find_unused_parameters_true'
else:
strategy = 'auto'
gpu_names = [torch.cuda.get_device_name(i) for i in devices]
print(f"Using GPU(s): {devices} -> {gpu_names}")
else:
accelerator = 'cpu'
devices = 1
strategy = 'auto'
num_gpus = 1
print("Using CPU (GPU not available or gpus=0)")
# 自动按实际使用的 GPU 数量调整全局 Batch Size 与学习率
config.TRAINER.WORLD_SIZE = num_gpus
config.TRAINER.TRUE_BATCH_SIZE = num_gpus * args.batch_size
_scaling = config.TRAINER.TRUE_BATCH_SIZE / config.TRAINER.CANONICAL_BS
config.TRAINER.TRUE_LR = config.TRAINER.CANONICAL_LR * _scaling
config.TRAINER.WARMUP_STEP = math.floor(config.TRAINER.WARMUP_STEP / _scaling)
print("Training Configuration:")
print(f" - Total Effective Batch Size: {config.TRAINER.TRUE_BATCH_SIZE}")
print(f" - Scaled Learning Rate: {config.TRAINER.TRUE_LR:.6f}")
print(f" - Warmup steps: {config.TRAINER.WARMUP_STEP}")
print("=" * 60)
# 初始化模型与加载权重
try:
model = PL_CRFT(config, pretrained_ckpt=args.pretrained_ckpt)
print("Model loaded with pretrained weights successfully.")
except Exception as e:
print(f"Failed to load pretrained weights: {e}")
model = PL_CRFT(config)
print("Model created with random initial weights.")
total_params = sum(p.numel() for p in model.parameters())
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f"Total model parameters: {total_params:,}")
print(f"Trainable parameters: {trainable_params:,}")
print("=" * 60)
# 创建数据模块
data_module = MultiSceneDataModule(args, config)
print("Data module created.")
# 设置 TensorBoard 日志路径与 Checkpoint 保存
logger_tb = TensorBoardLogger(save_dir='logs/tb_logs', name=args.exp_name)
ckpt_dir = Path(logger_tb.log_dir) / 'checkpoints'
callbacks = [
LearningRateMonitor(logging_interval='step'),
ModelCheckpoint(
monitor='val_AEPE',
save_top_k=5,
mode='min',
save_last=True,
dirpath=str(ckpt_dir),
filename='{epoch:02d}-{val_AEPE:.3f}'
)
]
# 实例化 Trainer
trainer = pl.Trainer(
max_epochs=args.epochs,
accelerator=accelerator,
devices=devices, # 传入 list 格式的设备列表,如 [0]、[1] 或 [0, 1]
strategy=strategy, # 动态配置 DDP 策略
logger=logger_tb,
callbacks=callbacks,
gradient_clip_val=1.0,
precision="32-true",
check_val_every_n_epoch=1,
log_every_n_steps=10 # 实时打日志
)
print("Trainer created.")
print(f"Logs will be saved to: {logger_tb.log_dir}")
print("=" * 60)
# 开始训练
try:
print("Starting training process...")
start_time = time.time()
trainer.fit(model, datamodule=data_module)
end_time = time.time()
total_time = end_time - start_time
hours, remainder = divmod(total_time, 3600)
minutes, seconds = divmod(remainder, 60)
print("\n" + "=" * 60)
print(f"Training completed in {int(hours)}h {int(minutes)}m {seconds:.2f}s")
avg_epoch_time = total_time / args.epochs
print(f"Average time per epoch: {avg_epoch_time:.2f}s")
print(f"Best checkpoints saved at: {ckpt_dir}")
print("=" * 60)
except Exception as e:
print(f"\nTraining failed with exception: {e}")
import traceback
traceback.print_exc()
if __name__ == '__main__':
main()