New ideas

This commit is contained in:
2025-11-21 07:16:05 +01:00
parent 1538b0acf7
commit 5f8493d59e
6 changed files with 85738 additions and 115 deletions
+21 -22
View File
@@ -17,7 +17,7 @@ class BaseDataset(Dataset):
n_windows_per_file,
step=1,
feature_columns=("lat", "lon", "alt", "ias"),
context_columns=(f"type_encoding_{i}" for i in range(4), "last_lat", "last_lon", "last_alt", "last_ias"),
context_columns=(*[f"type_encoding_{i}" for i in range(4)], "last_lat", "last_lon", "last_alt", "last_ias"),
time_columns=("timestamp",),
target_columns=("lat", "lon", "alt"),
device="cpu",
@@ -33,6 +33,7 @@ class BaseDataset(Dataset):
# Which columns to use
sample = pd.read_csv(self.filepaths[0], nrows=5)
sample = self._modify_df(sample)
self.feature_cols = list(feature_columns)
self.context_cols = list(context_columns)
self.time_cols = list(time_columns)
@@ -69,27 +70,36 @@ class BaseDataset(Dataset):
def _modify_df(self, df):
"""Hook for subclasses to modify dataframe before slicing windows."""
df.loc[:, ["lat", "lon", "alt"]] = preprocessors.norm_coords(
lat_norm, lon_norm, alt_norm = preprocessors.norm_coords(
df["lat"].values, df["lon"].values, df["alt"].values
)
df.loc[:, "lat"] = lat_norm
df.loc[:, "lon"] = lon_norm
df.loc[:, "alt"] = alt_norm
df.loc[:, "ias"] = preprocessors.norm_ias(df["ias"].values)
df.loc[:, "dt"] = df["timestamp"].diff()
df.loc[:, "dt"] = preprocessors.fillna_with_mean(df["dt"], allow_nan_mean=False)
df.loc[:, "timestamp"] = preprocessors.norm_time(df["timestamp"].values)
for col in ["lat", "lon", "alt", "ias"]:
df.loc[:, col] = preprocessors.fillna_with_mean(df[col])
if df[col].isnull().all() and col != "ias":
raise ValueError(f"All values in column {col} are NaN.")
df.loc[:, col] = preprocessors.fillna_with_mean(df[col], allow_nan_mean=(col == "ias"))
mapped_df = preprocessors.map_categories(df, preprocessors.category_mappings)
X_a_1 = F.one_hot(torch.tensor(mapped_df["R_1_IDX"].values)).float()
X_a_2 = F.one_hot(torch.tensor(mapped_df["R_2_IDX"].values)).float()
X_t = F.one_hot(torch.tensor(mapped_df["T_IDX"].values)).float()
max_values = list(preprocessors.category_max_limits.values())
X_a_1 = F.one_hot(torch.tensor(mapped_df["R_1_IDX"].values), max_values[0]).float()
X_a_2 = F.one_hot(torch.tensor(mapped_df["R_2_IDX"].values), max_values[1]).float()
X_t = F.one_hot(torch.tensor(mapped_df["T_IDX"].values), max_values[2]).float()
df.loc[:, [f"type_encoding_{i}" for i in range(4)]] = preprocessors.encode_features(
X_a_1, X_a_2, X_t
)
last_row = df.iloc[-1][["lat", "lon", "alt", "ias"]]
last_row = df.iloc[-1][["lat", "lon", "alt", "ias", "timestamp"]]
df.loc[:, "last_lat"] = last_row["lat"]
df.loc[:, "last_lon"] = last_row["lon"]
df.loc[:, "last_alt"] = last_row["alt"]
df.loc[:, "last_ias"] = last_row["ias"]
df.loc[:, "last_timestamp"] = last_row["timestamp"]
return df
@@ -103,7 +113,7 @@ class EvenlySpacedDataset(BaseDataset):
n_windows_per_file,
step=1,
feature_columns=("lat", "lon", "alt", "ias"),
context_columns=("r","t"),
context_columns=(*[f"type_encoding_{i}" for i in range(4)], "last_lat", "last_lon", "last_alt", "last_ias"),
time_columns=("timestamp",),
target_columns=("lat", "lon", "alt"),
device="cpu",
@@ -145,7 +155,7 @@ class EvenlySpacedDataset(BaseDataset):
X_feat = window.iloc[: self.n_input][self.feature_cols].to_numpy(np.float32)
if len(self.context_cols) > 0:
X_context = window.iloc[: self.n_input][self.context_cols].to_numpy(np.float32)
X_context = window.iloc[0][self.context_cols].to_numpy(np.float32)
else:
X_context = None
X_time = window.iloc[: self.window_size][self.time_cols].to_numpy(np.float32)
@@ -166,7 +176,7 @@ class EvenlySpacedStreamingDataset(BaseDataset):
n_windows_per_file,
step=1,
feature_columns=("lat", "lon", "alt", "ias"),
context_columns=("r","t"),
context_columns=(*[f"type_encoding_{i}" for i in range(4)], "last_lat", "last_lon", "last_alt", "last_ias"),
time_columns=("timestamp",),
target_columns=("lat", "lon", "alt"),
device="cpu",
@@ -232,7 +242,7 @@ def get_datasets(
n_windows_per_file: int,
step: int = 1,
feature_columns=("lat", "lon", "alt", "ias"),
context_columns=("r","t"),
context_columns=(*[f"type_encoding_{i}" for i in range(4)], "last_lat", "last_lon", "last_alt", "last_ias"),
time_columns=("timestamp",),
target_columns=("lat", "lon", "alt"),
device="cpu",
@@ -298,14 +308,3 @@ def get_datasets(
return train_dataset, val_dataset, test_dataset
class EncodingDataset(Dataset):
def __init__(self, X_a_1, X_a_2, X_b):
super().__init__()
self.X_a_1 = X_a_1
self.X_a_2 = X_a_2
self.X_b = X_b
def __len__(self):
return self.X_a_1.shape[0]
def __getitem__(self, idx):
return self.X_a_1[idx], self.X_a_2[idx], self.X_b[idx]