-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
111 lines (91 loc) · 4.07 KB
/
Copy pathtrain.py
File metadata and controls
111 lines (91 loc) · 4.07 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
"""联邦训练入口
用法:
python train.py --config configs/default.yaml --device cpu
python train.py --strategy fedavg --rounds 50 --clients 10
"""
import sys
import argparse
import torch
from pathlib import Path
sys.path.insert(0, str(Path(__file__).parent))
from src.config import Config
from src.data.dataset import create_dataloaders
from src.models.multimodal_model import MultimodalSentimentModel
from src.federated.client import FederatedClient
from src.federated.server import FederatedServer
def main():
parser = argparse.ArgumentParser(description="Fed-MSA 联邦训练")
parser.add_argument("--config", default="configs/default.yaml")
parser.add_argument("--device", default="cpu")
parser.add_argument("--strategy", default="fedavg")
parser.add_argument("--rounds", type=int, default=0)
parser.add_argument("--clients", type=int, default=0)
parser.add_argument("--missing-rate", type=float, default=0.3)
args = parser.parse_args()
# 配置
config = Config()
if Path(args.config).exists():
config = Config.from_yaml(args.config)
# 输出到块存储(不占系统盘)
output_root = Path("/root/blockdata/fed-msa")
(output_root / "models").mkdir(parents=True, exist_ok=True)
(output_root / "results").mkdir(parents=True, exist_ok=True)
if args.rounds > 0:
config.federated.rounds = args.rounds
if args.clients > 0:
config.federated.num_clients = args.clients
if args.strategy != "fedavg":
config.federated.strategy = args.strategy
device = torch.device(args.device)
print(f"Device: {device}, Strategy: {config.federated.strategy}")
print(f"Clients: {config.federated.num_clients}, Rounds: {config.federated.rounds}")
# 数据加载器 + 预先加载到内存
train_loader, val_loader, test_loader = create_dataloaders(
config.data.csd_dir, batch_size=config.federated.batch_size,
missing_rate=args.missing_rate,
max_seq_len=config.data.max_seq_len)
# 在线程池之前预先加载,避免多线程竞争
train_loader.dataset._preload()
val_loader.dataset._preload()
test_loader.dataset._preload()
# 全局模型
global_model = MultimodalSentimentModel(config)
print(f"Model params: {sum(p.numel() for p in global_model.parameters()):,}")
# 服务器
server = FederatedServer(global_model, strategy=config.federated.strategy,
device=device)
# 创建客户端 (多GPU轮询分配,共享预加载的数据集)
num_gpus = torch.cuda.device_count() if device.type == 'cuda' else 1
print(f"GPUs available: {num_gpus}")
clients = []
for i in range(config.federated.num_clients):
gpu_id = i % num_gpus
client_device = torch.device(f"cuda:{gpu_id}") if device.type == 'cuda' else device
client = FederatedClient(
client_id=i, model=global_model,
train_loader=train_loader, val_loader=val_loader,
learning_rate=config.federated.learning_rate,
weight_decay=config.federated.weight_decay,
local_epochs=config.federated.local_epochs,
device=client_device, strategy=config.federated.strategy,
)
clients.append(client)
print(f"{len(clients)} clients on {num_gpus} GPUs")
# 联邦训练循环
from tqdm import tqdm
for r in tqdm(range(config.federated.rounds), desc="FL Training"):
metrics = server.run_round(clients, test_loader)
if (r + 1) % 10 == 0 or r == 0:
tqdm.write(f"Round {r+1:3d}/{config.federated.rounds} | "
f"MAE={metrics.get('mae',0):.4f} | "
f"Corr={metrics.get('corr',0):.4f}")
# 保存
if (r + 1) % 10 == 0:
server.save_checkpoint(str(output_root / f"models/checkpoint_round_{r+1}.pt"))
# 最终保存
server.save_checkpoint(str(output_root / "models/final_model.pt"))
server.save_history(str(output_root / "results/training_history.json"))
print(f"\nBest Corr: {server.best_corr:.4f}")
print("Done.")
if __name__ == "__main__":
main()