Files
aiRtrafficNN/Code/python/notebooks/predict.py
T
2025-12-16 19:59:02 +01:00

103 lines
3.0 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
def load_file_list(path: pathlib.Path):
base = path.parent
return [
(base / f).resolve()
for f in path.read_text().splitlines()
if (base / f).exists()
]
file_list_path = pathlib.Path("../../cpp/known_routes_and_aircraft.csv")
file_list = load_file_list(file_list_path)
base_ds = dataloader.EvenlySpacedDataset(
filepaths=file_list,
n_input=warm,
n_output=pred,
n_windows_per_file=5,
step=step,
feature_columns=("lat", "lon", "alt", "ias"),
context_columns=("last_lat", "last_lon", "last_alt", "last_ias", "last_timestamp"),
time_columns=("timestamp", "dt"),
target_columns=("lat", "lon", "alt"),
)
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"))
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_pred = lat_pred.flatten()
lon_pred = lon_pred.flatten()
alt_pred = alt_pred.flatten()
lat_true, lon_true, alt_true = denorm_coords(y[...,0], y[...,1], y[...,2])
lat_true = lat_true.flatten()
lon_true = lon_true.flatten()
alt_true = alt_true.flatten()
results.append(torch.stack([lat_true, lon_true, alt_true, lat_pred, lon_pred, alt_pred], dim=1))
results_tensor = torch.cat(results, dim=0)
np_y_all = results_tensor.cpu().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="",
)