prop/docs/plans/2026-04-01-algorithm-refinement-ml-training.md
Graham McIntire 69b5caf876
Normalize ML features to prevent NaN gradient explosion
Raw features had vastly different scales (pressure ~1013, sin/cos ~[-1,1])
causing gradient explosion. Normalize all atmospheric features to ~[0,1]
using known physical bounds. Add Polaris dep for optimizer.
2026-04-01 09:08:33 -05:00

991 lines
32 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# Algorithm Refinement & ML Training Implementation Plan
> **For Claude:** REQUIRED SUB-SKILL: Use superpowers:executing-plans to implement this plan task-by-task.
**Goal:** Refine the propagation scoring algorithm using 58K QSOs matched to HRRR conditions, update algo.md and scorer code, then train the Axon ML model for fast HRRR-based prediction.
**Architecture:** Analysis mix task extracts QSO↔HRRR correlations → findings drive scorer/algo.md updates → recomputed scores become ML training data → trained model integrates into grid worker with algorithm fallback.
**Tech Stack:** Elixir/Ecto (analysis queries), Nx/Axon/EXLA (ML training), PostgreSQL (data source)
---
### Task 1: Analysis Mix Task — QSO-HRRR Dataset Builder
**Files:**
- Create: `lib/mix/tasks/propagation_analyze.ex`
- Test: `test/mix/tasks/propagation_analyze_test.exs`
**Step 1: Write the failing test**
```elixir
# test/mix/tasks/propagation_analyze_test.exs
defmodule Mix.Tasks.PropagationAnalyzeTest do
use Microwaveprop.DataCase
alias Mix.Tasks.PropagationAnalyze
describe "build_dataset/0" do
test "returns list of QSO-condition maps with required keys" do
# Insert a QSO with known position and timestamp
qso = insert_qso(
band: 10_000,
distance_km: 200.0,
qso_timestamp: ~N[2023-06-15 12:00:00],
pos1: %{"lat" => 33.0, "lng" => -97.0},
pos2: %{"lat" => 34.0, "lng" => -96.0}
)
# Insert HRRR profiles at the grid points nearest to both endpoints
insert_hrrr_profile(
lat: 33.0, lon: -97.0, valid_time: ~N[2023-06-15 12:00:00],
surface_temp_c: 30.0, surface_dewpoint_c: 20.0,
surface_pressure_mb: 1010.0, min_refractivity_gradient: -100.0,
hpbl_m: 500.0, pwat_mm: 30.0, surface_refractivity: 330.0,
ducting_detected: false
)
insert_hrrr_profile(
lat: 34.0, lon: -96.0, valid_time: ~N[2023-06-15 12:00:00],
surface_temp_c: 28.0, surface_dewpoint_c: 18.0,
surface_pressure_mb: 1012.0, min_refractivity_gradient: -80.0,
hpbl_m: 600.0, pwat_mm: 25.0, surface_refractivity: 320.0,
ducting_detected: false
)
[row] = PropagationAnalyze.build_dataset()
assert row.qso_id == qso.id
assert row.band == 10_000
assert row.distance_km == 200.0
assert row.month == 6
assert row.utc_hour == 12
assert_in_delta row.avg_surface_temp_c, 29.0, 0.01
assert_in_delta row.avg_min_refractivity_gradient, -90.0, 0.01
assert_in_delta row.avg_hpbl_m, 550.0, 0.01
end
test "skips QSOs without HRRR match at either endpoint" do
insert_qso(
band: 10_000, distance_km: 200.0,
qso_timestamp: ~N[2023-06-15 12:00:00],
pos1: %{"lat" => 33.0, "lng" => -97.0},
pos2: %{"lat" => 34.0, "lng" => -96.0}
)
# Only insert one endpoint's HRRR profile
insert_hrrr_profile(
lat: 33.0, lon: -97.0, valid_time: ~N[2023-06-15 12:00:00],
surface_temp_c: 30.0, surface_dewpoint_c: 20.0,
surface_pressure_mb: 1010.0, min_refractivity_gradient: -100.0,
hpbl_m: 500.0, pwat_mm: 30.0
)
assert PropagationAnalyze.build_dataset() == []
end
end
end
```
**Step 2: Run test to verify it fails**
Run: `mix test test/mix/tasks/propagation_analyze_test.exs --trace`
Expected: FAIL — module not defined
**Step 3: Write test helpers**
Add to `test/support/fixtures.ex` (or create if missing) — helper functions `insert_qso/1` and `insert_hrrr_profile/1` that insert records directly via Repo.
**Step 4: Write minimal implementation**
```elixir
# lib/mix/tasks/propagation_analyze.ex
defmodule Mix.Tasks.PropagationAnalyze do
@shortdoc "Analyze QSO-HRRR correlations for algorithm refinement"
@moduledoc """
Builds a dataset matching QSOs to HRRR atmospheric conditions at both
endpoints, then runs correlation analysis against QSO distance.
## Usage
mix propagation.analyze
"""
use Mix.Task
import Ecto.Query
alias Microwaveprop.Repo
@impl Mix.Task
def run(_args) do
Application.put_env(:microwaveprop, Oban, queues: false)
Mix.Task.run("app.start")
IO.puts("Building QSO-HRRR analysis dataset...")
dataset = build_dataset()
IO.puts("Matched #{length(dataset)} QSOs to HRRR conditions")
IO.puts("\n=== CORRELATION ANALYSIS ===\n")
run_correlation_analysis(dataset)
IO.puts("\n=== FACTOR ANALYSIS ===\n")
run_factor_analysis(dataset)
IO.puts("\n=== INTERACTION EFFECTS ===\n")
run_interaction_analysis(dataset)
IO.puts("\nDone. Use findings to update algo.md and scorer.ex")
end
@doc """
Builds the QSO-HRRR matched dataset.
For each QSO with valid positions and a timestamp within HRRR range (2016+),
finds the nearest HRRR grid point to each endpoint (rounded to 0.125°) at
the nearest hour. Averages atmospheric conditions across both endpoints.
Returns a list of maps with:
- qso_id, band, distance_km, month, utc_hour
- avg_surface_temp_c, avg_surface_dewpoint_c, avg_surface_pressure_mb
- avg_min_refractivity_gradient, avg_hpbl_m, avg_pwat_mm
- avg_surface_refractivity, ducting_either (boolean)
- terrain_verdict
"""
def build_dataset do
# Query uses a CTE to round QSO positions to HRRR grid and truncate timestamps,
# then joins to hrrr_profiles at both endpoints, averaging atmospheric values.
query = """
WITH qso_grid AS (
SELECT
q.id AS qso_id,
q.band::integer AS band,
q.distance_km,
EXTRACT(MONTH FROM q.qso_timestamp)::integer AS month,
EXTRACT(HOUR FROM q.qso_timestamp)::integer AS utc_hour,
ROUND(CAST((q.pos1->>'lat')::float * 8 AS numeric)) / 8 AS grid_lat1,
ROUND(CAST((q.pos1->>'lng')::float * 8 AS numeric)) / 8 AS grid_lng1,
ROUND(CAST((q.pos2->>'lat')::float * 8 AS numeric)) / 8 AS grid_lat2,
ROUND(CAST((q.pos2->>'lng')::float * 8 AS numeric)) / 8 AS grid_lng2,
date_trunc('hour', q.qso_timestamp) AS match_time
FROM qsos q
WHERE q.distance_km < 3000
AND q.distance_km > 0
AND q.qso_timestamp >= '2016-06-30'
AND q.pos1->>'lng' IS NOT NULL
AND q.pos2->>'lng' IS NOT NULL
)
SELECT
qg.qso_id, qg.band, qg.distance_km, qg.month, qg.utc_hour,
(COALESCE(h1.surface_temp_c, h2.surface_temp_c) +
COALESCE(h2.surface_temp_c, h1.surface_temp_c)) / 2.0 AS avg_surface_temp_c,
(COALESCE(h1.surface_dewpoint_c, h2.surface_dewpoint_c) +
COALESCE(h2.surface_dewpoint_c, h1.surface_dewpoint_c)) / 2.0 AS avg_surface_dewpoint_c,
(COALESCE(h1.surface_pressure_mb, h2.surface_pressure_mb) +
COALESCE(h2.surface_pressure_mb, h1.surface_pressure_mb)) / 2.0 AS avg_surface_pressure_mb,
(COALESCE(h1.min_refractivity_gradient, h2.min_refractivity_gradient) +
COALESCE(h2.min_refractivity_gradient, h1.min_refractivity_gradient)) / 2.0 AS avg_min_refractivity_gradient,
(COALESCE(h1.hpbl_m, h2.hpbl_m) +
COALESCE(h2.hpbl_m, h1.hpbl_m)) / 2.0 AS avg_hpbl_m,
(COALESCE(h1.pwat_mm, h2.pwat_mm) +
COALESCE(h2.pwat_mm, h1.pwat_mm)) / 2.0 AS avg_pwat_mm,
(COALESCE(h1.surface_refractivity, h2.surface_refractivity) +
COALESCE(h2.surface_refractivity, h1.surface_refractivity)) / 2.0 AS avg_surface_refractivity,
COALESCE(h1.ducting_detected, false) OR COALESCE(h2.ducting_detected, false) AS ducting_either,
tp.verdict AS terrain_verdict
FROM qso_grid qg
JOIN hrrr_profiles h1
ON h1.lat = qg.grid_lat1 AND h1.lon = qg.grid_lng1 AND h1.valid_time = qg.match_time
JOIN hrrr_profiles h2
ON h2.lat = qg.grid_lat2 AND h2.lon = qg.grid_lng2 AND h2.valid_time = qg.match_time
LEFT JOIN terrain_profiles tp ON tp.qso_id = qg.qso_id
ORDER BY qg.band, qg.distance_km DESC
"""
Repo.query!(query, [], timeout: 300_000).rows
|> Enum.map(fn row ->
[qso_id, band, distance_km, month, utc_hour,
avg_temp, avg_dewpoint, avg_pressure, avg_gradient,
avg_hpbl, avg_pwat, avg_refractivity, ducting_either,
terrain_verdict] = row
%{
qso_id: qso_id,
band: band,
distance_km: distance_km,
month: month,
utc_hour: utc_hour,
avg_surface_temp_c: avg_temp,
avg_surface_dewpoint_c: avg_dewpoint,
avg_surface_pressure_mb: avg_pressure,
avg_min_refractivity_gradient: avg_gradient,
avg_hpbl_m: avg_hpbl,
avg_pwat_mm: avg_pwat,
avg_surface_refractivity: avg_refractivity,
ducting_either: ducting_either,
terrain_verdict: terrain_verdict
}
end)
end
# --- Correlation analysis ---
defp run_correlation_analysis(dataset) do
bands = [10_000, 24_000, 47_000, 75_000]
for band <- bands do
band_data = Enum.filter(dataset, &(&1.band == band))
next_if_empty(band_data, band, fn ->
IO.puts("--- Band: #{div(band, 1000)} GHz (n=#{length(band_data)}) ---")
distances = Enum.map(band_data, & &1.distance_km)
variables = [
{:avg_surface_temp_c, "Surface Temp (C)"},
{:avg_surface_dewpoint_c, "Dewpoint (C)"},
{:avg_surface_pressure_mb, "Pressure (mb)"},
{:avg_min_refractivity_gradient, "Refr. Gradient (N/km)"},
{:avg_hpbl_m, "BL Depth (m)"},
{:avg_pwat_mm, "PWAT (mm)"},
{:avg_surface_refractivity, "Surface Refractivity"},
{:month, "Month"},
{:utc_hour, "UTC Hour"}
]
for {key, label} <- variables do
values = Enum.map(band_data, &Map.get(&1, key))
corr = spearman_correlation(values, distances)
bar = correlation_bar(corr)
IO.puts(" #{String.pad_trailing(label, 25)} r=#{format_corr(corr)} #{bar}")
end
IO.puts("")
end)
end
end
# --- Factor analysis (binned median distance) ---
defp run_factor_analysis(dataset) do
band_10g = Enum.filter(dataset, &(&1.band == 10_000))
band_24g = Enum.filter(dataset, &(&1.band == 24_000))
for {band_data, label} <- [{band_10g, "10 GHz"}, {band_24g, "24 GHz"}] do
next_if_empty(band_data, label, fn ->
IO.puts("--- #{label} Factor Breakpoints ---\n")
# Refractivity gradient bins
IO.puts("Refractivity Gradient → Median Distance:")
print_binned(band_data, :avg_min_refractivity_gradient, :distance_km,
[{-999, -200}, {-200, -150}, {-150, -100}, {-100, -75}, {-75, -55}, {-55, -40}, {-40, 0}])
# Temperature bins
IO.puts("\nSurface Temperature (C) → Median Distance:")
print_binned(band_data, :avg_surface_temp_c, :distance_km,
[{-20, 0}, {0, 10}, {10, 20}, {20, 25}, {25, 30}, {30, 35}, {35, 50}])
# HPBL bins
IO.puts("\nBoundary Layer Depth (m) → Median Distance:")
print_binned(band_data, :avg_hpbl_m, :distance_km,
[{0, 200}, {200, 500}, {500, 1000}, {1000, 1500}, {1500, 2000}, {2000, 5000}])
# Pressure bins
IO.puts("\nPressure (mb) → Median Distance:")
print_binned(band_data, :avg_surface_pressure_mb, :distance_km,
[{980, 1000}, {1000, 1010}, {1010, 1015}, {1015, 1020}, {1020, 1025}, {1025, 1040}])
# PWAT bins
IO.puts("\nPWAT (mm) → Median Distance:")
print_binned(band_data, :avg_pwat_mm, :distance_km,
[{0, 10}, {10, 20}, {20, 30}, {30, 40}, {40, 50}, {50, 80}])
# Time of day
IO.puts("\nUTC Hour → Median Distance:")
for hour_range <- [{0, 3}, {3, 6}, {6, 9}, {9, 12}, {12, 15}, {15, 18}, {18, 21}, {21, 24}] do
{lo, hi} = hour_range
in_bin = Enum.filter(band_data, fn r ->
r.utc_hour >= lo and r.utc_hour < hi
end)
print_bin_stats(" #{lo}-#{hi} UTC", in_bin, :distance_km)
end
# Month
IO.puts("\nMonth → Median Distance:")
for m <- 1..12 do
in_bin = Enum.filter(band_data, &(&1.month == m))
month_name = Enum.at(~w(Jan Feb Mar Apr May Jun Jul Aug Sep Oct Nov Dec), m - 1)
print_bin_stats(" #{month_name}", in_bin, :distance_km)
end
# Terrain verdict
IO.puts("\nTerrain Verdict → Median Distance:")
for verdict <- ["BLOCKED", "CLEAR", "FRESNEL_PARTIAL", nil] do
in_bin = Enum.filter(band_data, &(&1.terrain_verdict == verdict))
print_bin_stats(" #{verdict || "UNKNOWN"}", in_bin, :distance_km)
end
IO.puts("")
end)
end
end
# --- Interaction effects ---
defp run_interaction_analysis(dataset) do
band_10g = Enum.filter(dataset, &(&1.band == 10_000))
next_if_empty(band_10g, "10 GHz interactions", fn ->
IO.puts("--- 10 GHz: Gradient × Time of Day ---")
for hour_range <- [{0, 6}, {6, 12}, {12, 18}, {18, 24}] do
{lo, hi} = hour_range
time_data = Enum.filter(band_10g, fn r -> r.utc_hour >= lo and r.utc_hour < hi end)
strong = Enum.filter(time_data, fn r ->
r.avg_min_refractivity_gradient != nil and r.avg_min_refractivity_gradient < -100
end)
weak = Enum.filter(time_data, fn r ->
r.avg_min_refractivity_gradient != nil and r.avg_min_refractivity_gradient >= -100
end)
strong_med = median(Enum.map(strong, & &1.distance_km))
weak_med = median(Enum.map(weak, & &1.distance_km))
IO.puts(" #{lo}-#{hi} UTC: strong gradient (<-100) median=#{format_km(strong_med)} (n=#{length(strong)}), " <>
"weak (>=-100) median=#{format_km(weak_med)} (n=#{length(weak)})")
end
IO.puts("\n--- 10 GHz: HPBL × Month ---")
for month_group <- [{[6, 7, 8], "Summer"}, {[12, 1, 2], "Winter"}, {[3, 4, 5], "Spring"}, {[9, 10, 11], "Fall"}] do
{months, label} = month_group
seasonal = Enum.filter(band_10g, fn r -> r.month in months end)
shallow = Enum.filter(seasonal, fn r -> r.avg_hpbl_m != nil and r.avg_hpbl_m < 500 end)
deep = Enum.filter(seasonal, fn r -> r.avg_hpbl_m != nil and r.avg_hpbl_m >= 1000 end)
shallow_med = median(Enum.map(shallow, & &1.distance_km))
deep_med = median(Enum.map(deep, & &1.distance_km))
IO.puts(" #{label}: shallow BL (<500m) median=#{format_km(shallow_med)} (n=#{length(shallow)}), " <>
"deep BL (>=1000m) median=#{format_km(deep_med)} (n=#{length(deep)})")
end
IO.puts("\n--- 10 GHz: Ducting Detection vs Distance ---")
ducting = Enum.filter(band_10g, & &1.ducting_either)
no_ducting = Enum.filter(band_10g, &(!&1.ducting_either))
IO.puts(" Ducting detected: median=#{format_km(median(Enum.map(ducting, & &1.distance_km)))} (n=#{length(ducting)})")
IO.puts(" No ducting: median=#{format_km(median(Enum.map(no_ducting, & &1.distance_km)))} (n=#{length(no_ducting)})")
end)
end
# --- Statistical helpers ---
defp spearman_correlation(xs, ys) do
xs_clean = Enum.zip(xs, ys) |> Enum.reject(fn {x, _} -> is_nil(x) end)
if length(xs_clean) < 10, do: nil, else: do_spearman(xs_clean)
end
defp do_spearman(pairs) do
n = length(pairs)
{xs, ys} = Enum.unzip(pairs)
rx = rank(xs)
ry = rank(ys)
d_sq = Enum.zip(rx, ry) |> Enum.map(fn {a, b} -> (a - b) * (a - b) end) |> Enum.sum()
1 - 6 * d_sq / (n * (n * n - 1))
end
defp rank(values) do
indexed = values |> Enum.with_index() |> Enum.sort_by(fn {v, _i} -> v end)
ranked = indexed |> Enum.with_index(1) |> Enum.map(fn {{_v, orig_i}, rank} -> {orig_i, rank} end)
ranked |> Enum.sort_by(fn {orig_i, _rank} -> orig_i end) |> Enum.map(fn {_i, r} -> r * 1.0 end)
end
defp median([]), do: nil
defp median(list) do
sorted = Enum.sort(list)
n = length(sorted)
mid = div(n, 2)
if rem(n, 2) == 0 do
(Enum.at(sorted, mid - 1) + Enum.at(sorted, mid)) / 2
else
Enum.at(sorted, mid)
end
end
defp percentile([], _p), do: nil
defp percentile(list, p) do
sorted = Enum.sort(list)
k = (length(sorted) - 1) * p / 100
f = trunc(k)
c = Float.ceil(k) |> trunc()
if f == c, do: Enum.at(sorted, f), else: Enum.at(sorted, f) + (k - f) * (Enum.at(sorted, c) - Enum.at(sorted, f))
end
defp print_binned(data, group_key, value_key, bins) do
for {lo, hi} <- bins do
in_bin = Enum.filter(data, fn row ->
v = Map.get(row, group_key)
v != nil and v >= lo and v < hi
end)
label = " [#{lo}, #{hi})"
print_bin_stats(String.pad_trailing(label, 18), in_bin, value_key)
end
end
defp print_bin_stats(label, data, value_key) do
n = length(data)
if n > 0 do
values = Enum.map(data, &Map.get(&1, value_key))
med = median(values)
p25 = percentile(values, 25)
p75 = percentile(values, 75)
IO.puts("#{label} n=#{String.pad_leading("#{n}", 6)} median=#{format_km(med)} p25=#{format_km(p25)} p75=#{format_km(p75)}")
else
IO.puts("#{label} n= 0")
end
end
defp next_if_empty([], label, _fun), do: IO.puts("Skipping #{label}: no data")
defp next_if_empty(_data, _label, fun), do: fun.()
defp format_corr(nil), do: " N/A"
defp format_corr(corr), do: :io_lib.format("~7.4f", [corr]) |> to_string()
defp format_km(nil), do: " N/A"
defp format_km(km), do: :io_lib.format("~7.1f km", [km]) |> to_string()
defp correlation_bar(nil), do: ""
defp correlation_bar(corr) do
len = round(abs(corr) * 20)
bar = String.duplicate(if(corr >= 0, do: "+", else: "-"), len)
"|#{bar}|"
end
end
```
**Step 5: Run test to verify it passes**
Run: `mix test test/mix/tasks/propagation_analyze_test.exs --trace`
Expected: PASS
**Step 6: Run the analysis against real data**
Run: `mix propagation.analyze 2>&1 | tee docs/analysis_output.txt`
This will take several minutes (joining 58K QSOs against 54M HRRR profiles). Save the full output for reference during algo.md updates.
**Step 7: Commit**
```bash
git add lib/mix/tasks/propagation_analyze.ex test/mix/tasks/propagation_analyze_test.exs
git commit -m "Add mix propagation.analyze task for QSO-HRRR correlation analysis"
```
---
### Task 2: Interpret Analysis & Update algo.md
**Files:**
- Modify: `algo.md`
**Step 1: Read the analysis output**
Read `docs/analysis_output.txt` and identify:
- Which factors have strong vs weak correlation with distance
- Where current scoring thresholds diverge from data-suggested breakpoints
- Any interaction effects that the current independent-factor model misses
- Whether any current factors should be removed or new ones added
**Step 2: Update algo.md Part 2 (Key Empirical Findings)**
Add a new section "Data-Driven Refinements (April 2026)" documenting:
- Per-factor correlation strengths by band
- Revised threshold curves with data justification
- Weight adjustments with before/after comparison
- Any added or removed factors
**Step 3: Update algo.md Part 4 (Scoring Functions)**
Revise each factor's documentation to match new thresholds and behavior.
**Step 4: Update algo.md weights table**
Show old vs new weights with justification.
**Step 5: Commit**
```bash
git add algo.md
git commit -m "Update algo.md with data-driven findings from 58K QSO analysis"
```
---
### Task 3: Update Scorer & BandConfig
**Files:**
- Modify: `lib/microwaveprop/propagation/scorer.ex`
- Modify: `lib/microwaveprop/propagation/band_config.ex`
- Modify: `test/microwaveprop/propagation/scorer_test.exs`
- Modify: `test/microwaveprop/propagation/band_config_test.exs`
**Step 1: Update band_config.ex weights**
Change `@weights` map to match the data-derived optimal weights from the analysis.
**Step 2: Update band_config.ex thresholds**
Update `@refractivity_thresholds`, `@humidity_beneficial_thresholds`, seasonal tables, and any other thresholds that the analysis showed should change.
**Step 3: Update scorer.ex scoring functions**
Revise scoring function breakpoints to match the data-derived natural breakpoints. If the analysis warrants adding or removing factors, modify `composite_score/2` and the weights map accordingly.
**Step 4: Update tests to match new behavior**
Update `scorer_test.exs` and `band_config_test.exs` assertions to match the new thresholds and weights. Keep the same test structure — just adjust expected values.
**Step 5: Run tests**
Run: `mix test test/microwaveprop/propagation/scorer_test.exs test/microwaveprop/propagation/band_config_test.exs --trace`
Expected: PASS
**Step 6: Run full test suite**
Run: `mix test`
Expected: All pass
**Step 7: Commit**
```bash
git add lib/microwaveprop/propagation/scorer.ex lib/microwaveprop/propagation/band_config.ex
git add test/microwaveprop/propagation/scorer_test.exs test/microwaveprop/propagation/band_config_test.exs
git commit -m "Update scorer thresholds and weights from data analysis"
```
---
### Task 4: Training Mix Task
**Files:**
- Create: `lib/mix/tasks/propagation_train.ex`
- Modify: `lib/microwaveprop/propagation/model.ex` (add train/2, evaluate/2)
- Test: `test/mix/tasks/propagation_train_test.exs`
- Modify: `test/microwaveprop/propagation/model_test.exs`
**Step 1: Write failing test for Model.train/2**
```elixir
# In test/microwaveprop/propagation/model_test.exs, add:
describe "train/2" do
test "trains model and returns params with reduced loss" do
# Generate synthetic training data: 100 samples
features = Nx.random_uniform({100, 13}, type: :f32)
targets = Nx.random_uniform({100, 1}, type: :f32)
{params, metrics} = Model.train(features, targets, epochs: 5, batch_size: 32)
# Params should have the right structure
assert Map.has_key?(params, "hidden_1")
assert Map.has_key?(params, "hidden_2")
assert Map.has_key?(params, "output")
# Should return training metrics
assert is_float(metrics.final_loss)
assert metrics.final_loss >= 0
end
end
describe "evaluate/2" do
test "returns RMSE and R-squared on test data" do
params = Model.init()
features = Nx.random_uniform({50, 13}, type: :f32)
targets = Nx.random_uniform({50, 1}, type: :f32)
metrics = Model.evaluate(params, features, targets)
assert is_float(metrics.rmse)
assert is_float(metrics.r_squared)
assert metrics.rmse >= 0
end
end
```
**Step 2: Run test to verify it fails**
Run: `mix test test/microwaveprop/propagation/model_test.exs --trace`
Expected: FAIL — train/2 and evaluate/2 not defined
**Step 3: Implement Model.train/2 and Model.evaluate/2**
Add to `lib/microwaveprop/propagation/model.ex`:
```elixir
@doc """
Trains the model on the given features and targets tensors.
Options:
- :epochs — number of training epochs (default 50)
- :batch_size — samples per batch (default 256)
- :learning_rate — Adam learning rate (default 0.001)
- :patience — early stopping patience (default 5)
Returns {trained_params, %{final_loss: float}}.
"""
def train(features, targets, opts \\ []) do
epochs = Keyword.get(opts, :epochs, 50)
batch_size = Keyword.get(opts, :batch_size, 256)
lr = Keyword.get(opts, :learning_rate, 0.001)
model = build()
loop =
model
|> Axon.Loop.trainer(:mean_squared_error, Polaris.Optimizers.adam(learning_rate: lr))
|> Axon.Loop.metric(:mean_absolute_error)
data =
features
|> Nx.to_batched(batch_size)
|> Stream.zip(Nx.to_batched(targets, batch_size))
|> Stream.map(fn {x, y} -> {%{"features" => x}, y} end)
params = Axon.Loop.run(loop, data, Axon.ModelState.empty(),
epochs: epochs,
compiler: EXLA
)
# Compute final loss
{_init, predict_fn} = Axon.build(model)
preds = predict_fn.(params, %{"features" => features})
final_loss = Nx.mean(Nx.pow(Nx.subtract(preds, targets), 2)) |> Nx.to_number()
{params, %{final_loss: final_loss}}
end
@doc """
Evaluates model on test data. Returns %{rmse: float, r_squared: float}.
"""
def evaluate(params, features, targets) do
model = build()
{_init, predict_fn} = Axon.build(model)
preds = predict_fn.(params, %{"features" => features})
mse = Nx.mean(Nx.pow(Nx.subtract(preds, targets), 2)) |> Nx.to_number()
rmse = :math.sqrt(mse)
# R-squared
ss_res = Nx.sum(Nx.pow(Nx.subtract(targets, preds), 2)) |> Nx.to_number()
mean_y = Nx.mean(targets) |> Nx.to_number()
ss_tot = Nx.sum(Nx.pow(Nx.subtract(targets, mean_y), 2)) |> Nx.to_number()
r_squared = if ss_tot > 0, do: 1.0 - ss_res / ss_tot, else: 0.0
%{rmse: rmse, r_squared: r_squared}
end
```
**Step 4: Run model tests**
Run: `mix test test/microwaveprop/propagation/model_test.exs --trace`
Expected: PASS
**Step 5: Write the training mix task**
```elixir
# lib/mix/tasks/propagation_train.ex
defmodule Mix.Tasks.PropagationTrain do
@shortdoc "Train ML model from propagation scores + HRRR data"
@moduledoc """
Builds training data from propagation_scores joined to hrrr_profiles,
trains the Axon model, evaluates on held-out test set, and saves weights.
## Usage
mix propagation.train
mix propagation.train --epochs 100 --batch-size 512
"""
use Mix.Task
alias Microwaveprop.Propagation.Model
alias Microwaveprop.Repo
@impl Mix.Task
def run(args) do
{opts, _, _} = OptionParser.parse(args,
strict: [epochs: :integer, batch_size: :integer, sample: :integer])
Application.put_env(:microwaveprop, Oban, queues: false)
Mix.Task.run("app.start")
epochs = Keyword.get(opts, :epochs, 50)
batch_size = Keyword.get(opts, :batch_size, 256)
sample_size = Keyword.get(opts, :sample, 500_000)
IO.puts("Loading training data (sample=#{sample_size})...")
{features, targets, band_counts} = load_training_data(sample_size)
n = Nx.axis_size(features, 0)
IO.puts("Loaded #{n} samples")
for {band, count} <- Enum.sort(band_counts), do: IO.puts(" #{div(band, 1000)} GHz: #{count}")
# 80/10/10 split
n_train = round(n * 0.8)
n_val = round(n * 0.1)
train_x = Nx.slice(features, [0, 0], [n_train, 13])
train_y = Nx.slice(targets, [0, 0], [n_train, 1])
val_x = Nx.slice(features, [n_train, 0], [n_val, 13])
val_y = Nx.slice(targets, [n_train, 0], [n_val, 1])
test_x = Nx.slice(features, [n_train + n_val, 0], [n - n_train - n_val, 13])
test_y = Nx.slice(targets, [n_train + n_val, 0], [n - n_train - n_val, 1])
IO.puts("\nSplit: train=#{n_train}, val=#{n_val}, test=#{n - n_train - n_val}")
IO.puts("Training for #{epochs} epochs, batch_size=#{batch_size}...\n")
{params, train_metrics} = Model.train(train_x, train_y,
epochs: epochs, batch_size: batch_size)
IO.puts("\nTraining complete. Final loss: #{Float.round(train_metrics.final_loss, 6)}")
IO.puts("\n=== Evaluation ===")
val_metrics = Model.evaluate(params, val_x, val_y)
IO.puts("Validation: RMSE=#{Float.round(val_metrics.rmse * 100, 2)} points, R²=#{Float.round(val_metrics.r_squared, 4)}")
test_metrics = Model.evaluate(params, test_x, test_y)
IO.puts("Test: RMSE=#{Float.round(test_metrics.rmse * 100, 2)} points, R²=#{Float.round(test_metrics.r_squared, 4)}")
IO.puts("\nSaving model to priv/models/propagation_v1.nx...")
Model.save(params)
IO.puts("Done!")
end
defp load_training_data(sample_size) do
query = """
SELECT
h.surface_temp_c,
h.surface_dewpoint_c,
h.surface_pressure_mb,
h.min_refractivity_gradient,
h.hpbl_m,
h.pwat_mm,
EXTRACT(HOUR FROM ps.valid_time)::integer AS utc_hour,
EXTRACT(MONTH FROM ps.valid_time)::integer AS month,
ps.band_mhz,
ps.score
FROM propagation_scores ps
JOIN hrrr_profiles h
ON h.lat = ps.lat AND h.lon = ps.lon AND h.valid_time = ps.valid_time
WHERE h.surface_temp_c IS NOT NULL
AND h.surface_dewpoint_c IS NOT NULL
AND h.min_refractivity_gradient IS NOT NULL
ORDER BY RANDOM()
LIMIT $1
"""
result = Repo.query!(query, [sample_size], timeout: 600_000)
band_counts = result.rows
|> Enum.frequencies_by(fn row -> Enum.at(row, 8) end)
{feature_rows, target_rows} =
result.rows
|> Enum.map(fn [temp, dewpoint, pressure, gradient, hpbl, pwat, hour, month, band_mhz, score] ->
features = Model.encode_features(%{
surface_temp_c: temp,
surface_dewpoint_c: dewpoint,
surface_pressure_mb: pressure,
min_refractivity_gradient: gradient,
hpbl_m: hpbl || 500.0,
pwat_mm: pwat || 20.0,
utc_hour: hour,
month: month,
freq_mhz: band_mhz
})
target = score / 100.0
{features, [target]}
end)
|> Enum.unzip()
features = Nx.tensor(feature_rows, type: :f32)
targets = Nx.tensor(target_rows, type: :f32)
{features, targets, band_counts}
end
end
```
**Step 6: Run the training**
Run: `mix propagation.train --epochs 50 --sample 500000 2>&1 | tee docs/training_output.txt`
**Step 7: Commit**
```bash
git add lib/mix/tasks/propagation_train.ex lib/microwaveprop/propagation/model.ex
git add test/mix/tasks/propagation_train_test.exs test/microwaveprop/propagation/model_test.exs
git commit -m "Add ML training pipeline with Axon for propagation prediction"
```
---
### Task 5: Integrate ML Model into Grid Worker
**Files:**
- Modify: `lib/microwaveprop/propagation.ex` (add ML prediction path)
- Modify: `lib/microwaveprop/propagation/model.ex` (add predict_score/2)
- Modify: `test/microwaveprop/propagation_test.exs`
- Modify: `test/microwaveprop/propagation/model_test.exs`
**Step 1: Write failing test for Model.predict_score/2**
```elixir
# In model_test.exs
describe "predict_score/2" do
test "returns integer score 0-100 from conditions map" do
params = Model.init() # random weights, but should still return valid range
conditions = %{
surface_temp_c: 25.0,
surface_dewpoint_c: 15.0,
surface_pressure_mb: 1013.0,
min_refractivity_gradient: -80.0,
hpbl_m: 500.0,
pwat_mm: 25.0,
utc_hour: 12,
month: 6,
freq_mhz: 10_000
}
score = Model.predict_score(params, conditions)
assert is_integer(score)
assert score >= 0 and score <= 100
end
end
```
**Step 2: Run test to verify it fails**
Run: `mix test test/microwaveprop/propagation/model_test.exs --trace`
Expected: FAIL — predict_score/2 not defined
**Step 3: Implement Model.predict_score/2**
```elixir
@doc """
Predicts a 0-100 propagation score from a conditions map.
Convenience wrapper around encode_features + predict.
"""
def predict_score(params, conditions) do
features =
conditions
|> encode_features()
|> Nx.tensor(type: :f32)
|> Nx.reshape({1, @feature_count})
params
|> predict(features)
|> Nx.squeeze()
|> Nx.multiply(100)
|> Nx.round()
|> Nx.to_number()
|> trunc()
|> max(0)
|> min(100)
end
```
**Step 4: Run test**
Run: `mix test test/microwaveprop/propagation/model_test.exs --trace`
Expected: PASS
**Step 5: Add ML path to Propagation.score_grid_point_with_data/5**
Modify `lib/microwaveprop/propagation.ex` to try ML prediction first, fall back to algorithm:
```elixir
# In score_grid_point_with_data/5, after building conditions map:
# Try ML model if loaded, fall back to algorithm scorer
for band_config <- BandConfig.all_bands() do
result = case ml_params() do
nil ->
Scorer.composite_score(conditions, band_config)
params ->
ml_conditions = Map.merge(conditions, %{freq_mhz: band_config.freq_mhz})
score = Model.predict_score(params, ml_conditions)
%{score: score, factors: %{ml_predicted: true}}
end
Map.put(result, :band_mhz, band_config.freq_mhz)
end
```
With a module-level cached model load:
```elixir
@ml_params_key :propagation_ml_params
def load_ml_model do
case Model.load() do
{:ok, params} ->
:persistent_term.put(@ml_params_key, params)
:ok
:error ->
:ok
end
end
defp ml_params do
try do
:persistent_term.get(@ml_params_key)
rescue
ArgumentError -> nil
end
end
```
**Step 6: Update Propagation tests**
Add tests verifying:
- Without ML model loaded, score_grid_point uses algorithm (existing behavior)
- With ML model loaded, score_grid_point uses ML predictions
**Step 7: Run full test suite**
Run: `mix test`
Expected: All pass
**Step 8: Commit**
```bash
git add lib/microwaveprop/propagation.ex lib/microwaveprop/propagation/model.ex
git add test/microwaveprop/propagation_test.exs test/microwaveprop/propagation/model_test.exs
git commit -m "Integrate ML model into grid worker with algorithm fallback"
```
---
### Task 6: Load ML Model at Application Startup
**Files:**
- Modify: `lib/microwaveprop/application.ex`
**Step 1: Add ML model loading to application start**
After Repo and other services start, call `Propagation.load_ml_model()` to cache params in persistent_term. This is a no-op if the model file doesn't exist.
**Step 2: Run full test suite**
Run: `mix test`
Expected: All pass
**Step 3: Run precommit**
Run: `mix precommit`
Expected: All pass
**Step 4: Commit**
```bash
git add lib/microwaveprop/application.ex
git commit -m "Load ML model at application startup"
```