79 lines
2.3 KiB
Python
79 lines
2.3 KiB
Python
import torch
|
|
import pathlib
|
|
import pandas as pd
|
|
from aiRNN import dataloader, models, losses
|
|
from aiRNN.preprocessors import denorm_coords
|
|
import numpy as np
|
|
import copy
|
|
|
|
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
|
|
|
|
|
|
def split_dataset(dataset, frac=0.8, seed=42):
|
|
n = len(dataset)
|
|
train_n = int(frac * n)
|
|
return torch.utils.data.random_split(
|
|
dataset,
|
|
[train_n, n - train_n],
|
|
generator=torch.Generator().manual_seed(seed),
|
|
)
|
|
|
|
step = 10
|
|
base_name = "LSTM"
|
|
cls = models.ThreeInputLSTM
|
|
hidden_size = 16
|
|
rnn_size = 64
|
|
hidden_layers = 0
|
|
rnn_layers = 3
|
|
rnn_dropout = 0.0
|
|
altitude_weight = 1e-3
|
|
warm = 900
|
|
pred = 1600
|
|
start_offset = 30 * 60 - warm
|
|
end_offset = 30 * 60 - pred
|
|
base_ds = dataloader.SaveDataset(
|
|
torch.load("long_dataset.pt"),
|
|
step=step,
|
|
start_offset=start_offset,
|
|
end_offset=end_offset,
|
|
)
|
|
train_ds, val_ds = split_dataset(base_ds)
|
|
model = cls(
|
|
time_in=2,
|
|
feat_in=4,
|
|
context_in=5,
|
|
hidden_size=hidden_size,
|
|
rnn_size=rnn_size,
|
|
out_size=3,
|
|
hidden_layers=hidden_layers,
|
|
rnn_layers=rnn_layers,
|
|
rnn_dropout=rnn_dropout,
|
|
device=DEVICE,
|
|
)
|
|
|
|
model.load_state_dict(torch.load("LSTM_wu900_ps150_aw0.001.pt", map_location=torch.device('cpu')))
|
|
model.to(DEVICE)
|
|
model.eval()
|
|
with torch.no_grad():
|
|
val_loader = torch.utils.data.DataLoader(val_ds, batch_size=64, collate_fn=dataloader.collate_to_cpu, num_workers=4)
|
|
results = []
|
|
for X_f, X_t, y, X_c in val_loader:
|
|
X_f = X_f.to(DEVICE)
|
|
X_t = X_t.to(DEVICE)
|
|
y = y.to(DEVICE)
|
|
if X_c is not None:
|
|
X_c = X_c.to(DEVICE)
|
|
y_pred, _ = model(X_t, X_f, X_c, warm // step, pred // step)
|
|
lat_pred, lon_pred, alt_pred = denorm_coords(y_pred[...,0], y_pred[...,1], y_pred[...,2])
|
|
lat_true, lon_true, alt_true = denorm_coords(y[...,0], y[...,1], y[...,2])
|
|
results.append(torch.stack([lat_true, lon_true, alt_true, lat_pred, lon_pred, alt_pred], dim=1).cpu())
|
|
|
|
np_y_all = torch.cat(results, dim=0).numpy()
|
|
np.savetxt(
|
|
f"predictions_{base_name}_wu{warm}_ps{pred}_aw{altitude_weight}.csv",
|
|
np_y_all,
|
|
delimiter=",",
|
|
header="true_lat,true_lon,true_alt,pred_lat,pred_lon,pred_alt",
|
|
comments="",
|
|
)
|