| """Predict the test set and perform hour/season-window weather normalization.""" |
|
|
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import yaml |
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
| from model.meteonorm_rf import BASE_FEATURES, encode_features, load_checkpoint |
|
|
|
|
| def circular_day_distance(a, b): |
| distance = np.abs(a - b) |
| return np.minimum(distance, 365 - distance) |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| data = np.load(ROOT / config["data"]["path"]) |
| model, metadata = load_checkpoint(ROOT / config["paths"]["checkpoint"]) |
| x = encode_features(data) |
| test = np.asarray(metadata["test_indices"], dtype=np.int64) |
| observed = data["pollution"][test] |
| predicted = model.predict(x[test]) |
| repeats = int(config["normalization"]["engineering_resamples"]) |
| window = int(config["normalization"]["season_window_days"]) |
| rng = np.random.default_rng(int(config["seed"]) + 17) |
| normalized = np.empty_like(predicted) |
| weather_columns = [BASE_FEATURES.index(name) for name in |
| ("wind_speed", "wind_direction", "pressure", "temperature", "relative_humidity")] |
| all_hour, all_doy = data["hour"], data["day_of_year"] |
| for output_index, row_index in enumerate(test): |
| candidates = np.flatnonzero((all_hour == all_hour[row_index]) & |
| (circular_day_distance(all_doy, all_doy[row_index]) <= window)) |
| draws = rng.choice(candidates, repeats, replace=True) |
| replicated = np.repeat(x[row_index:row_index + 1], repeats, axis=0) |
| replicated[:, weather_columns] = x[draws][:, weather_columns] |
| |
| normalized[output_index] = model.predict(replicated).mean(0) |
| output = ROOT / config["paths"]["inference"] |
| output.parent.mkdir(parents=True, exist_ok=True) |
| np.savez_compressed(output, test_indices=test, observed=observed, predicted=predicted, |
| normalized=normalized, ttrend=data["ttrend"][test], |
| day_of_year=data["day_of_year"][test], hour=data["hour"][test], |
| station_id=data["station_id"][test], pollutant_names=data["pollutant_names"], |
| resamples=np.array(repeats), season_window_days=np.array(window)) |
| print(f"predictions={output.relative_to(ROOT)} test_rows={len(test)} normalization_resamples={repeats}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|