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_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.cat([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="", )