monaaaaaa's picture
Upload 26 files
eeabcff verified
Raw
History Blame Contribute Delete
10.6 kB
from utils import *
from modules import *
import os, sys
import numpy as np
from tqdm import tqdm
import random
import torch
from torch import nn
from config import CFG
from dataset import *
import torch.utils.data
import copy, json, pickle
import itertools as it
def make_next_record_dir(basedir, prefix=''):
path = '%s/%%s001/' % basedir
n = 2
while os.path.exists(path % prefix):
path = '%s/%%s%.3d/' % (basedir, n)
n += 1
pth = path % prefix
os.makedirs(pth)
return pth
def setup_seed(seed):
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
np.random.seed(seed)
random.seed(seed)
torch.backends.cudnn.deterministic = True
def my_collate(batch):
batch = list(filter(lambda x: (x is not None), batch))
msbinl, molfpl, molfml, vl, al, msl = [], [], [], [], [], []
bat = {}
msbinl1, msbinl2, msbinl3 = [], [], []
for b in batch:
if 'ms_bins' in b:
msbinl.append(b['ms_bins'])
if 'ms_bins1' in b:
msbinl1.append(b['ms_bins1'])
if 'ms_bins2' in b:
msbinl2.append(b['ms_bins2'])
if 'ms_bins3' in b and b['ms_bins3'] is not None:
msbinl3.append(b['ms_bins3'])
if 'mol_fps' in b:
molfpl.append(b['mol_fps'])
if 'mol_fmvec' in b:
molfml.append(b['mol_fmvec'])
if 'V' in b:
vl.append(b['V'])
if 'A' in b:
al.append(b['A'])
if 'mol_size' in b:
msl.append(b['mol_size'])
if msbinl:
bat['ms_bins'] = torch.stack(msbinl)
if msbinl1:
bat['ms_bins1'] = torch.stack(msbinl1)
if msbinl2:
bat['ms_bins2'] = torch.stack(msbinl2)
if msbinl3:
bat['ms_bins3'] = torch.stack(msbinl3)
if molfpl:
bat['mol_fps'] = torch.stack(molfpl)
if molfml:
bat['mol_fmvec'] = torch.stack(molfml)
if vl and al and msl:
max_n = max(map(lambda x: x.shape[0], vl))
vl1, al1 = [], []
for v in vl:
vl1.append(pad_V(v, max_n))
for a in al:
al1.append(pad_A(a, max_n))
bat['V'] = torch.stack(vl1)
bat['A'] = torch.stack(al1)
bat['mol_size'] = torch.cat(msl, dim=0)
# return torch.utils.data.dataloader.default_collate(batch)
return bat
def make_train_valid(data, valid_ratio, seed=1234):
idxs = np.arange(len(data))
np.random.seed(seed)
np.random.shuffle(idxs)
lenval = int(valid_ratio * len(data))
valid_set = [data[i] for i in idxs[:lenval]]
train_set = [data[i] for i in idxs[lenval:]]
return train_set, valid_set
def build_loaders(inp, mode, cfg, num_workers):
if type(inp[0]) is dict:
dataset = Dataset(inp, cfg)
else:
dataset = PathDataset(inp, cfg)
dataloader = torch.utils.data.DataLoader(
dataset,
batch_size=cfg.batch_size,
num_workers=num_workers,
shuffle=True if mode == "train" else False,
collate_fn=my_collate
)
return dataloader
def train_epoch(model, train_loader, optimizer, lr_scheduler, step):
loss_meter = AvgMeter()
tqdm_object = tqdm(train_loader, total=len(train_loader))
for batch in tqdm_object:
for k, v in batch.items():
batch[k] = v.to(CFG.device)
loss = model(batch)
optimizer.zero_grad()
loss.backward()
optimizer.step()
if step == "batch":
lr_scheduler.step()
count = batch["ms_bins"].size(0)
loss_meter.update(loss.item(), count)
tqdm_object.set_postfix(train_loss=loss_meter.avg, lr=get_lr(optimizer))
return loss_meter
def valid_epoch(model, valid_loader):
loss_meter = AvgMeter()
tqdm_object = tqdm(valid_loader, total=len(valid_loader))
for batch in tqdm_object:
for k, v in batch.items():
batch[k] = v.to(CFG.device)
loss = model(batch)
count = batch["ms_bins"].size(0)
loss_meter.update(loss.item(), count)
tqdm_object.set_postfix(valid_loss=loss_meter.avg)
return loss_meter
def main(data, cfg=CFG, savedir='data/train', encmodel=None, ratio=1):
setup_seed(cfg.seed)
train_data_file = cfg.train_data_file
valid_data_file = cfg.valid_data_file
if train_data_file.endswith('.pt'):
train_set = torch.load(train_data_file)
elif train_data_file.endswith('.json'):
train_set = json.load(open(train_data_file, 'r', encoding='utf-8'))
if valid_data_file.endswith('.pt'):
valid_set = torch.load(valid_data_file)
elif valid_data_file.endswith('.json'):
valid_set = json.load(open(valid_data_file, 'r', encoding='utf-8'))
if os.path.isdir(train_data_file):
train_set = []
for i in tqdm(range(cfg.train_number_data), desc='load train ...'):
tmp_file = train_data_file + str(i) + ".pt"
if os.path.exists(tmp_file):
# tmp_d = torch.load(tmp_file)
# train_set.append(tmp_d)
train_set.append(tmp_file)
print("len train data ...", len(train_set))
print("len valid_set data ...", len(valid_set))
# train_data_file = "data/train_data.json"
# valid_data_file = "data/valid_data.json"
#
# if os.path.exists(train_data_file):
# train_set = json.load(open(train_data_file, 'r', encoding='utf-8'))
# valid_set = json.load(open(valid_data_file, 'r', encoding='utf-8'))
# else:
# train_set, valid_set = make_train_valid(data, valid_ratio=cfg.valid_ratio, seed=cfg.seed)
#
# json.dump(train_set, open(train_data_file, 'w', encoding='utf-8'))
# json.dump(valid_set, open(valid_data_file, 'w', encoding='utf-8'))
#
# n = len(train_set)
# if ratio < 1:
# train_set = random.sample(train_set, int(n*ratio))
# print(f'Ratio {ratio}, lenall {n}, newtrainset {len(train_set)}')
train_loader = build_loaders(train_set, "train", cfg, 10)
valid_loader = build_loaders(valid_set, "valid", cfg, 10)
step = "epoch"
best_loss = float('inf')
best_model_fn = ''
best_model_fns = []
# model = FragSimiModel(cfg).to(cfg.device)
model = FragSimiModelNew(cfg).to(cfg.device)
if not encmodel is None:
model.mol_gnn_encoder.load_state_dict(encmodel.mol_gnn_encoder.state_dict())
# fraze mol_gnn_encoder weights
'''for name, param in model.named_parameters():
if 'mol_gnn_encoder' in name:
print(152, 'fraze mol_gnn_encoder weights')
param.requires_grad = False'''
print(model)
print(cfg.device)
print(model.feature_proj.bias.device)
optimizer = torch.optim.AdamW(
model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay
)
lr_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
optimizer, mode="min", patience=cfg.patience, factor=cfg.factor
)
for epoch in range(cfg.epochs):
print(f"Epoch: {epoch + 1}/{cfg.epochs}")
model.train()
train_loss = train_epoch(model, train_loader, optimizer, lr_scheduler, step)
model.eval()
with torch.no_grad():
valid_loss = valid_epoch(model, valid_loader)
if True: # valid_loss.avg < best_loss:
best_loss = valid_loss.avg
best_model_fn = f"{savedir}/model-tloss{round(train_loss.avg, 3)}-vloss{round(valid_loss.avg, 3)}-epoch{epoch}.pth"
best_model_fn_base = best_model_fn.replace('.pth', '')
n = 1
while os.path.exists(best_model_fn):
best_model_fn = best_model_fn_base + f'-{n}.pth'
n += 1
checkpoint = {'state_dict': model.state_dict(), 'optimizer': optimizer.state_dict(), 'config': dict(CFG)}
best_model_fns.append(best_model_fn)
torch.save(checkpoint, best_model_fn)
print("Saved Best Model!")
best_model_fnl = []
for fn in best_model_fns:
if os.path.exists(fn):
best_model_fnl.append(fn)
for fn in best_model_fnl[:-cfg.keep_best_models_num]:
os.remove(fn)
best_model_fnl = best_model_fnl[-cfg.keep_best_models_num:]
print(best_model_fnl, best_loss)
return best_model_fnl, best_loss
if __name__ == "__main__":
# try:
# conffn = sys.argv[1]
# if conffn.endswith('.json'):
# CFG.load(sys.argv[1])
# elif conffn.endswith('.pth'):
# dpath = CFG.dataset_path
# d = torch.load(conffn)
# CFG.load(d['config'])
# CFG.dataset_path = dpath
# print('Use config from', conffn)
# except:
# pass
#
# try:
# savedir = sys.argv[2]
# except:
# savedir = 'out_data/'
# os.system('mkdir -p %s' % savedir)
savedir = 'out_data/'
mg = None
print(CFG)
if os.path.isdir(CFG.dataset_path):
# data = [os.path.join(CFG.dataset_path, i) for i in os.listdir(CFG.dataset_path) if i.endswith('mgf')]
# elif os.path.isfile(CFG.dataset_path):
# if CFG.dataset_path.endswith('.pkl'):
# data = pickle.load(open(CFG.dataset_path, 'rb'))
# else:
# data = json.load(open(CFG.dataset_path))
# pklfn = CFG.dataset_path.replace('.json', '.pkl')
# if not os.path.exists(pklfn):
# pickle.dump(data, open(pklfn, 'wb'))
# 支持多级目录通配符
data_files = []
for root, _, files in os.walk(CFG.dataset_path):
for f in files:
if f.endswith(('.json', '.pkl', '.mgf')):
data_files.append(os.path.join(root, f))
data = data_files
elif '*' in CFG.dataset_path: # 新增通配符支持
import glob
data = glob.glob(CFG.dataset_path)
elif os.path.isfile(CFG.dataset_path):
data = [CFG.dataset_path]
subdir = make_next_record_dir(savedir, f'train-')
os.system(f'cp -a *py {subdir}; cp -a GNN {subdir}')
CFG.save(f'{subdir}/config.json')
modelfnl, _ = main(data, CFG, subdir, mg)
# CUDA_VISIBLE_DEVICES=7 nohup python3 -u train.py > train_024_14.log 2>&1 &