BlidReview's picture
weights, code, eval script
bdce880 verified
Raw
History Blame Contribute Delete
8.31 kB
import os
import argparse
import matplotlib.pyplot as plt
import numpy as np
import torch
from tqdm import *
from utils.testloss import TestLoss
from model_dict import get_model
from utils.normalizer import UnitTransformer
parser = argparse.ArgumentParser('Training Transformer')
parser.add_argument('--lr', type=float, default=1e-3)
parser.add_argument('--epochs', type=int, default=500)
parser.add_argument('--weight_decay', type=float, default=1e-5)
parser.add_argument('--model', type=str, default='Transolver_1D')
parser.add_argument('--n-hidden', type=int, default=64, help='hidden dim')
parser.add_argument('--n-layers', type=int, default=3, help='layers')
parser.add_argument('--n-heads', type=int, default=4)
parser.add_argument('--batch-size', type=int, default=8)
parser.add_argument("--gpu", type=str, default='1', help="GPU index to use")
parser.add_argument('--max_grad_norm', type=float, default=None)
parser.add_argument('--downsample', type=int, default=5)
parser.add_argument('--mlp_ratio', type=int, default=1)
parser.add_argument('--dropout', type=float, default=0.0)
parser.add_argument('--ntrain', type=int, default=1000)
parser.add_argument('--unified_pos', type=int, default=0)
parser.add_argument('--ref', type=int, default=8)
parser.add_argument('--slice_num', type=int, default=32)
parser.add_argument('--eval', type=int, default=0)
parser.add_argument('--save_name', type=str, default='elas_Transolver')
parser.add_argument('--data_path', type=str, default='/data/fno')
args = parser.parse_args()
eval = args.eval
save_name = args.save_name
os.environ["CUDA_VISIBLE_DEVICES"] = args.gpu
def count_parameters(model):
total_params = 0
for name, parameter in model.named_parameters():
if not parameter.requires_grad: continue
params = parameter.numel()
total_params += params
print(f"Total Trainable Params: {total_params}")
return total_params
def main():
ntrain = args.ntrain
ntest = 200
PATH_Sigma = args.data_path + '/elasticity/Meshes/Random_UnitCell_sigma_10.npy'
PATH_XY = args.data_path + '/elasticity/Meshes/Random_UnitCell_XY_10.npy'
input_s = np.load(PATH_Sigma)
input_s = torch.tensor(input_s, dtype=torch.float).permute(1, 0)
input_xy = np.load(PATH_XY)
input_xy = torch.tensor(input_xy, dtype=torch.float).permute(2, 0, 1)
train_s = input_s[:ntrain]
test_s = input_s[-ntest:]
train_xy = input_xy[:ntrain]
test_xy = input_xy[-ntest:]
print(input_s.shape, input_xy.shape)
y_normalizer = UnitTransformer(train_s)
train_s = y_normalizer.encode(train_s)
y_normalizer.cuda()
train_loader = torch.utils.data.DataLoader(torch.utils.data.TensorDataset(train_xy, train_xy, train_s),
batch_size=args.batch_size,
shuffle=True)
test_loader = torch.utils.data.DataLoader(torch.utils.data.TensorDataset(test_xy, test_xy, test_s),
batch_size=args.batch_size,
shuffle=False)
print("Dataloading is over.")
model = get_model(args).Model(space_dim=2,
n_layers=args.n_layers,
n_hidden=args.n_hidden,
dropout=args.dropout,
n_head=args.n_heads,
Time_Input=False,
mlp_ratio=args.mlp_ratio,
fun_dim=0,
out_dim=1,
slice_num=args.slice_num,
ref=args.ref,
unified_pos=args.unified_pos).cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
print(args)
print(model)
count_parameters(model)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs)
myloss = TestLoss(size_average=False)
if eval:
model.load_state_dict(torch.load("./checkpoints/" + save_name + ".pt"))
model.eval()
if not os.path.exists('./results/' + save_name + '/'):
os.makedirs('./results/' + save_name + '/')
rel_err = 0.0
showcase = 10
id = 0
with torch.no_grad():
for pos, fx, y in test_loader:
id += 1
x, fx, y = pos.cuda(), fx.cuda(), y.cuda()
out = model(x, None).squeeze(-1)
out = y_normalizer.decode(out)
tl = myloss(out, y).item()
rel_err += tl
if id < showcase:
print(id)
plt.axis('off')
plt.scatter(x=fx[0, :, 0].detach().cpu().numpy(), y=fx[0, :, 1].detach().cpu().numpy(),
c=y[0, :].detach().cpu().numpy(), cmap='coolwarm')
plt.colorbar()
plt.clim(0, 1000)
plt.savefig(
os.path.join('./results/' + save_name + '/',
"gt_" + str(id) + ".pdf"), bbox_inches='tight', pad_inches=0)
plt.close()
plt.axis('off')
plt.scatter(x=fx[0, :, 0].detach().cpu().numpy(), y=fx[0, :, 1].detach().cpu().numpy(),
c=out[0, :].detach().cpu().numpy(), cmap='coolwarm')
plt.colorbar()
plt.clim(0, 1000)
plt.savefig(
os.path.join('./results/' + save_name + '/',
"pred_" + str(id) + ".pdf"), bbox_inches='tight', pad_inches=0)
plt.close()
plt.axis('off')
plt.scatter(x=fx[0, :, 0].detach().cpu().numpy(), y=fx[0, :, 1].detach().cpu().numpy(),
c=((y[0, :] - out[0, :])).detach().cpu().numpy(), cmap='coolwarm')
plt.clim(-8, 8)
plt.colorbar()
plt.savefig(
os.path.join('./results/' + save_name + '/',
"error_" + str(id) + ".pdf"), bbox_inches='tight', pad_inches=0)
plt.close()
rel_err /= ntest
print("rel_err : {}".format(rel_err))
else:
for ep in range(args.epochs):
model.train()
train_loss = 0
for pos, fx, y in train_loader:
x, fx, y = pos.cuda(), fx.cuda(), y.cuda() # x:B,N,2 fx:B,N,2 y:B,N,
optimizer.zero_grad()
out = model(x, None).squeeze(-1)
out = y_normalizer.decode(out)
y = y_normalizer.decode(y)
loss = myloss(out, y)
loss.backward()
if args.max_grad_norm is not None:
torch.nn.utils.clip_grad_norm_(model.parameters(), args.max_grad_norm)
optimizer.step()
train_loss += loss.item()
scheduler.step()
train_loss = train_loss / ntrain
print("Epoch {} Train loss : {:.5f}".format(ep, train_loss))
model.eval()
rel_err = 0.0
with torch.no_grad():
for pos, fx, y in test_loader:
x, fx, y = pos.cuda(), fx.cuda(), y.cuda()
out = model(x, None).squeeze(-1)
out = y_normalizer.decode(out)
tl = myloss(out, y).item()
rel_err += tl
rel_err /= ntest
print("rel_err : {}".format(rel_err))
if ep % 100 == 0:
if not os.path.exists('./checkpoints'):
os.makedirs('./checkpoints')
print('save model')
torch.save(model.state_dict(), os.path.join('./checkpoints', save_name + '.pt'))
if not os.path.exists('./checkpoints'):
os.makedirs('./checkpoints')
print('save model')
torch.save(model.state_dict(), os.path.join('./checkpoints', save_name + '.pt'))
if __name__ == "__main__":
main()