332 lines
11 KiB
Python
332 lines
11 KiB
Python
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)
|
||
|