MOS2-seg/code/train_pl.py

332 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import os
import json
import glob
import shutil
import torch
import argparse
torch.set_float32_matmul_precision('high')
import numpy as np
from sklearn.utils import shuffle
from torch import nn
from torch.utils.data import DataLoader
from torch.utils.data.sampler import *
from torch.optim.lr_scheduler import ReduceLROnPlateau
from pytorch_lightning.loggers import TensorBoardLogger
try:
from aim.pytorch_lightning import AimLogger
AIM_AVAILABLE = True
except ImportError:
AIM_AVAILABLE = False
print("Warning: Aim is not installed. Install it with: pip install aim")
from pytorch_lightning.callbacks import ModelCheckpoint
from pytorch_lightning.callbacks.early_stopping import EarlyStopping
import pytorch_lightning as pl
from core.model import *
from core.data import *
from core.metrics import *
# 在一个目录下找第一个.pth或.ckpt模型文件返回绝对路径
import uuid
def find_first_pth(dir_path):
import os
import glob
# 递归搜索 .pth 文件
pth_list = glob.glob(os.path.join(dir_path, '**', '*.pth'), recursive=True)
if len(pth_list) == 0:
return None
else:
return pth_list[0]
def extract_zip_files(source_dir):
import zipfile
from pathlib import Path
# 创建一个临时目录
dir_path = '/tmp-' + uuid.uuid4().hex
if not os.path.exists(dir_path):
os.makedirs(dir_path)
print("临时目录为:", dir_path)
# 遍历指定目录下的所有文件和子目录
for item in Path(source_dir).rglob('*.zip'):
# 确保是文件而不是目录
print("dataset item is :", item)
if item.is_file():
# 解压ZIP文件到临时目录
with zipfile.ZipFile(item, 'r') as zip_ref:
zip_ref.extractall(dir_path)
# 返回临时目录的路径
return dir_path
def get_args():
parser = argparse.ArgumentParser(description="Training configuration")
# Training params
parser.add_argument("--batch_size", type=int, default=64, help="Batch size")
parser.add_argument("--epochs", type=int, default=1000, help="Number of training epochs")
parser.add_argument("--lr", type=float, default=3e-4, help="Learning rate")
parser.add_argument("--weight_decay", type=float, default=1e-5, help="Weight decay")
parser.add_argument("--rf", type=float, default=0.9, help="Reduction factor / ratio factor")
parser.add_argument("--num_workers", type=int, default=4, help="Number of dataloader workers")
# Model params
parser.add_argument("--dim", type=int, default=256, help="Feature dimension")
parser.add_argument("--num_classes", type=int, default=3, help="Number of classes")
parser.add_argument("--model_type", type=str, default="unet_resnet34", help="Model name")
parser.add_argument(
"--model_name",
type=str,
default="./ckpt/resnet34-333f7ec4.pth",
help="Pretrained checkpoint path"
)
# Callback / save params
parser.add_argument("--save_top_k", type=int, default=1, help="Save top-k checkpoints")
parser.add_argument("--early_stop", type=int, default=2, help="Early stopping patience")
parser.add_argument("--every_n_epochs", type=int, default=1, help="Checkpoint save frequency")
parser.add_argument("--model_output", type=str, default="./logs/", help="Log save path")
# Data params
parser.add_argument("--dataset", type=str, default="../data/", help="Image dataset path")
args = parser.parse_args()
#args.model_name = find_first_pth(args.model_name)
args.dataset = extract_zip_files(args.dataset)
print("train args is ", args)
return args
class SegModel(pl.LightningModule):
def __init__(self, args):
super().__init__()
self.save_hyperparameters()
self.args = args
self.model = SMPModelFactory(
model=args.model_type,
encoder_weights_path=args.model_name,
classes=args.num_classes
).get_model()
self.loss = CustomLoss()
def forward(self, x):
out = self.model(x)
return out
def configure_optimizers(self):
optimizer = torch.optim.Adam(self.parameters(), lr=self.args.lr, weight_decay=self.args.weight_decay)
scheduler = ReduceLROnPlateau(optimizer, factor=self.args.rf, mode='max', patience=2, min_lr=0)
return {
'optimizer': optimizer,
'lr_scheduler': scheduler,
'monitor': 'val_iou'
}
def training_step(self, train_batch, batch_idx):
x, y = train_batch
pd = self.model(x)
loss = self.loss(pd, y)
train_iou = iou(pd, y)
train_acc = accuracy(pd, y)
# 记录训练指标
self.log('train_loss', loss, on_epoch=True, prog_bar=True, logger=True)
self.log('train_iou', train_iou, on_epoch=True, prog_bar=True, logger=True)
self.log('train_acc', train_acc, on_epoch=True, prog_bar=True, logger=True)
return loss
def validation_step(self, val_batch, batch_idx):
x, y = val_batch
pd = self.model(x)
loss = self.loss(pd, y)
val_iou = iou(pd, y)
val_acc = accuracy(pd, y)
# 记录验证指标
self.log('val_loss', loss, on_epoch=True, prog_bar=True, logger=True)
self.log('val_iou', val_iou, on_epoch=True, prog_bar=True, logger=True)
self.log('val_acc', val_acc, on_epoch=True, prog_bar=True, logger=True)
def predict_step(self, batch, batch_idx):
x, lbl = batch
pred = self.model(x)
return pred, lbl
os.environ['EXPERIMENT_REMOTE_REPO'] = 'aim://192.168.60.11:30058'
# main
if __name__ == '__main__':
args = get_args()
train_img_list = np.array(glob.glob('{}/{}/*.jpg'.format(args.dataset, 'train'))).tolist()
valid_img_list = np.array(glob.glob('{}/{}/*.jpg'.format(args.dataset, 'valid'))).tolist()
print('Train nums: {}, Valid nums: {}.'.format(len(train_img_list), len(valid_img_list)))
train_dataset = CustomDataset(train_img_list, dim=args.dim, data_type='train')
valid_dataset = CustomDataset(valid_img_list, dim=args.dim, data_type='valid')
train_loader = DataLoader(
dataset=train_dataset,
batch_size=args.batch_size,
num_workers=args.num_workers,
drop_last=True,
)
valid_loader = DataLoader(
dataset=valid_dataset,
batch_size=args.batch_size,
shuffle=False,
num_workers=args.num_workers,
)
model = SegModel(args)
# 确保 model_output 目录存在,并使用绝对路径
model_output_dir = os.path.abspath(args.model_output)
os.makedirs(model_output_dir, exist_ok=True)
# 创建 loggers 列表
loggers = []
# 创建 TensorBoard logger
tb_logger = TensorBoardLogger(
name='',
save_dir=model_output_dir,
version=0,
default_hp_metric=False
)
loggers.append(tb_logger)
# 创建 Aim logger如果可用
if AIM_AVAILABLE:
# 从环境变量读取 Aim repo 地址,如果不存在则使用本地路径
aim_repo = os.environ.get('EXPERIMENT_REMOTE_REPO', os.path.join(model_output_dir, '.aim'))
print(f"Aim repo: {aim_repo}")
# 从环境变量读取 EXPERIMENT_RUN_ID
experiment_run_id = os.environ.get('EXPERIMENT_RUN_ID')
if experiment_run_id:
print(f"EXPERIMENT_RUN_ID: {experiment_run_id}")
# 准备超参数字典
hparams = {
'batch_size': args.batch_size,
'epochs': args.epochs,
'lr': args.lr,
'weight_decay': args.weight_decay,
'rf': args.rf,
'num_workers': args.num_workers,
'dim': args.dim,
'num_classes': args.num_classes,
'model_type': args.model_type,
'save_top_k': args.save_top_k,
'early_stop': args.early_stop,
'every_n_epochs': args.every_n_epochs,
}
# 如果存在 EXPERIMENT_RUN_ID添加到超参数中
if experiment_run_id:
hparams['id'] = experiment_run_id
aim_logger = AimLogger(
repo=aim_repo, # 从环境变量读取或使用本地路径
experiment='MOS2-train', # 实验名称
system_tracking_interval=None # 禁用系统 CPU 和内存指标记录
)
# 记录超参数
aim_logger.log_hyperparams(hparams)
loggers.append(aim_logger)
print("Aim logger initialized successfully")
else:
print("Aim logger not available, using TensorBoard only")
# 为了兼容性,保留 logger 变量指向第一个 logger
logger = loggers[0]
checkpoint_callback = ModelCheckpoint(
dirpath=model_output_dir,
save_top_k=1,
monitor='val_iou',
mode='max',
filename='model',
save_last=False
)
earlystop_callback = EarlyStopping(
monitor="val_iou",
mode="max",
min_delta=0.00,
patience=args.early_stop,
)
# training
trainer = pl.Trainer(
accelerator='gpu',
devices=1,
max_epochs=args.epochs,
logger=loggers, # 使用 loggers 列表,支持多个 logger
callbacks=[checkpoint_callback, earlystop_callback]
)
# 强制设置 TensorBoard logger 和 checkpoint 的路径,确保直接保存在 model_output 目录下,不创建 version_0 子目录
if isinstance(logger, TensorBoardLogger):
logger._log_dir = model_output_dir
logger._version = None
checkpoint_callback.dirpath = model_output_dir
checkpoint_callback.filename = 'model'
trainer.fit(
model,
train_loader,
valid_loader
)
# # 将 version_0 目录下的文件移动到 model_output 目录,然后删除 version_0 目录
# version_dir = os.path.join(model_output_dir, 'version_0')
# if os.path.exists(version_dir):
# # 移动 version_0 目录下的所有文件到 model_output 目录
# for item in os.listdir(version_dir):
# src = os.path.join(version_dir, item)
# dst = os.path.join(model_output_dir, item)
# if os.path.isdir(src):
# if os.path.exists(dst):
# shutil.rmtree(dst)
# shutil.move(src, dst)
# else:
# if os.path.exists(dst):
# os.remove(dst)
# shutil.move(src, dst)
# # 删除空的 version_0 目录
# try:
# os.rmdir(version_dir)
# except:
# pass
# inference
predictions = trainer.predict(
model=model,
dataloaders=valid_loader,
ckpt_path='best',
weights_only=False
)
preds = torch.squeeze(torch.concat([item[0] for item in predictions])).numpy().tolist()
labels = torch.squeeze(torch.concat([item[1] for item in predictions])).numpy().tolist()
results = {
'img_path': valid_img_list,
'pred': preds,
'label': labels,
}
results_json = json.dumps(results)
with open(os.path.join(trainer.log_dir, 'valid.json'), 'w+') as f:
f.write(results_json)
shutil.rmtree(os.path.join(model_output_dir, 'version_0'), ignore_errors=True)