Retrained on 60k month-balanced HRRR profiles through 2026-04-05 (test RMSE 1.7 score points, R² 0.97). Added explain_prediction/4 returning ranked feature contributions via batched finite-difference attribution, so the UI can show users which weather factors drove each prediction.
442 lines
14 KiB
Elixir
442 lines
14 KiB
Elixir
defmodule Microwaveprop.Propagation.Model do
|
||
@moduledoc """
|
||
Neural network model for microwave propagation prediction.
|
||
|
||
Uses Axon to define a model that predicts propagation conditions
|
||
(score 0-100 per band) from atmospheric and temporal features.
|
||
|
||
## Features (inputs)
|
||
|
||
Atmospheric (from HRRR):
|
||
- surface_temp_c — surface temperature
|
||
- surface_dewpoint_c — surface dewpoint
|
||
- surface_pressure_mb — surface pressure
|
||
- abs_humidity — absolute humidity (g/m³), derived
|
||
- td_depression — temp minus dewpoint (°C)
|
||
- min_refractivity_gradient — minimum dN/dh from profile
|
||
- hpbl_m — planetary boundary layer height
|
||
- pwat_mm — precipitable water
|
||
|
||
Temporal:
|
||
- utc_hour — hour of day (0-23), sin/cos encoded
|
||
- month — month of year (1-12), sin/cos encoded
|
||
|
||
Band:
|
||
- freq_mhz — operating frequency, log-scaled
|
||
|
||
## Target (output)
|
||
|
||
- score — propagation score (0-100), normalized to 0-1 for training
|
||
|
||
## Architecture
|
||
|
||
Feed-forward neural network with:
|
||
- Input: 20 features (12 atmospheric + 2 solar + 1 geographic + 4 cyclical temporal + 1 frequency)
|
||
- Hidden: 2 dense layers with ReLU activation
|
||
- Output: 1 sigmoid unit (score 0-1, scaled to 0-100)
|
||
|
||
The model is intentionally simple to start. We can add complexity
|
||
(LSTM for temporal sequences, attention, ensemble) once we have
|
||
baseline performance metrics.
|
||
"""
|
||
|
||
@feature_count 20
|
||
@models_dir Path.join(:code.priv_dir(:microwaveprop), "models")
|
||
@default_path Path.join(@models_dir, "propagation_v1.nx")
|
||
|
||
def default_path, do: @default_path
|
||
|
||
@doc """
|
||
Builds the Axon model graph. Does not initialize parameters.
|
||
|
||
Returns an `%Axon{}` struct ready for `Axon.build/2` or `Axon.Loop.trainer/3`.
|
||
"""
|
||
def build do
|
||
"features"
|
||
|> Axon.input(shape: {nil, @feature_count})
|
||
|> Axon.dense(128, activation: :relu, name: "hidden_1")
|
||
|> Axon.dropout(rate: 0.2, name: "dropout_1")
|
||
|> Axon.dense(64, activation: :relu, name: "hidden_2")
|
||
|> Axon.dropout(rate: 0.15, name: "dropout_2")
|
||
|> Axon.dense(32, activation: :relu, name: "hidden_3")
|
||
|> Axon.dropout(rate: 0.1, name: "dropout_3")
|
||
|> Axon.dense(1, activation: :sigmoid, name: "output")
|
||
end
|
||
|
||
@doc """
|
||
Initializes random model parameters.
|
||
|
||
Returns `%{...}` map of layer name → tensor params.
|
||
"""
|
||
def init do
|
||
model = build()
|
||
{init_fn, _predict_fn} = Axon.build(model)
|
||
template = Nx.template({1, @feature_count}, :f32)
|
||
init_fn.(template, Axon.ModelState.empty())
|
||
end
|
||
|
||
@doc """
|
||
Saves model parameters to disk. Defaults to `priv/models/propagation_v1.nx`.
|
||
"""
|
||
def save(params, path \\ @default_path) do
|
||
File.mkdir_p!(Path.dirname(path))
|
||
binary = Nx.serialize(params)
|
||
File.write!(path, binary)
|
||
:ok
|
||
end
|
||
|
||
@doc """
|
||
Loads model parameters from disk. Returns `{:ok, params}` or `:error`.
|
||
"""
|
||
def load(path \\ @default_path) do
|
||
case File.read(path) do
|
||
{:ok, binary} -> {:ok, Nx.deserialize(binary)}
|
||
{:error, _} -> :error
|
||
end
|
||
end
|
||
|
||
@doc """
|
||
Loads saved parameters if available, otherwise initializes random ones.
|
||
"""
|
||
def load_or_init(path \\ @default_path) do
|
||
case load(path) do
|
||
{:ok, params} -> params
|
||
:error -> init()
|
||
end
|
||
end
|
||
|
||
@doc """
|
||
Trains the model on the given features and targets.
|
||
|
||
Features should be an `{n, #{@feature_count}}` tensor and targets an `{n, 1}` tensor
|
||
with values in [0, 1].
|
||
|
||
Returns `{trained_state, %{final_loss: float}}`.
|
||
|
||
## Options
|
||
|
||
* `:epochs` - number of training epochs (default: 50)
|
||
* `:batch_size` - batch size (default: 256)
|
||
* `:learning_rate` - Adam learning rate (default: 0.001)
|
||
* `:initial_state` - pre-trained model state to resume from (default: empty)
|
||
"""
|
||
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)
|
||
initial_state = Keyword.get(opts, :initial_state)
|
||
|
||
model = build()
|
||
|
||
loop = Axon.Loop.trainer(model, :mean_squared_error, Polaris.Optimizers.adam(learning_rate: lr))
|
||
|
||
data =
|
||
features
|
||
|> Nx.to_batched(batch_size)
|
||
|> Stream.zip(Nx.to_batched(targets, batch_size))
|
||
|> Stream.map(fn {x, y} -> {%{"features" => x}, y} end)
|
||
|
||
init_state = initial_state || Axon.ModelState.empty()
|
||
|
||
trained_state =
|
||
Axon.Loop.run(loop, data, init_state,
|
||
epochs: epochs,
|
||
compiler: EXLA
|
||
)
|
||
|
||
# Compute final loss on full dataset
|
||
{_init_fn, predict_fn} = Axon.build(model, compiler: EXLA)
|
||
preds = predict_fn.(trained_state, %{"features" => features})
|
||
final_loss = preds |> Nx.subtract(targets) |> Nx.pow(2) |> Nx.mean() |> Nx.to_number()
|
||
|
||
{trained_state, %{final_loss: final_loss}}
|
||
end
|
||
|
||
@doc """
|
||
Evaluates trained parameters against features and targets.
|
||
|
||
Returns `%{rmse: float, r_squared: float}`.
|
||
"""
|
||
def evaluate(params, features, targets) do
|
||
model = build()
|
||
{_init_fn, predict_fn} = Axon.build(model, compiler: EXLA)
|
||
preds = predict_fn.(params, %{"features" => features})
|
||
|
||
mse = targets |> Nx.subtract(preds) |> Nx.pow(2) |> Nx.mean() |> Nx.to_number()
|
||
rmse = :math.sqrt(mse)
|
||
|
||
ss_res = targets |> Nx.subtract(preds) |> Nx.pow(2) |> Nx.sum() |> Nx.to_number()
|
||
mean_y = targets |> Nx.mean() |> Nx.to_number()
|
||
ss_tot = targets |> Nx.subtract(mean_y) |> Nx.pow(2) |> Nx.sum() |> 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
|
||
|
||
@doc """
|
||
Predicts a propagation score (0-100) from raw conditions.
|
||
|
||
Takes a compiled predict function, model parameters, and a conditions map.
|
||
Returns an integer score clamped to [0, 100].
|
||
"""
|
||
def predict_score(predict_fn, params, conditions) do
|
||
features =
|
||
conditions
|
||
|> encode_features()
|
||
|> Nx.tensor(type: :f32)
|
||
|> Nx.reshape({1, @feature_count})
|
||
|
||
params
|
||
|> predict_fn.(%{"features" => features})
|
||
|> Nx.squeeze()
|
||
|> Nx.multiply(100)
|
||
|> Nx.round()
|
||
|> Nx.to_number()
|
||
|> trunc()
|
||
|> max(0)
|
||
|> min(100)
|
||
end
|
||
|
||
@doc """
|
||
Predicts scores for multiple condition maps in a single batched forward pass.
|
||
|
||
Takes a compiled predict function, model parameters, and a list of conditions maps.
|
||
Returns a list of integer scores (0-100).
|
||
"""
|
||
@batch_chunk_size 10_000
|
||
|
||
def predict_scores_batch(predict_fn, params, conditions_list) do
|
||
conditions_list
|
||
|> Enum.chunk_every(@batch_chunk_size)
|
||
|> Enum.flat_map(fn chunk ->
|
||
feature_rows = Enum.map(chunk, &encode_features/1)
|
||
batch = Nx.tensor(feature_rows, type: :f32)
|
||
|
||
params
|
||
|> predict_fn.(%{"features" => batch})
|
||
|> Nx.multiply(100)
|
||
|> Nx.round()
|
||
|> Nx.squeeze(axes: [1])
|
||
|> Nx.to_flat_list()
|
||
|> Enum.map(fn score -> score |> trunc() |> max(0) |> min(100) end)
|
||
end)
|
||
end
|
||
|
||
@doc """
|
||
Compiles the predict function for reuse. Call once at load time.
|
||
"""
|
||
def compile_predict do
|
||
model = build()
|
||
{_init_fn, predict_fn} = Axon.build(model, compiler: EXLA)
|
||
predict_fn
|
||
end
|
||
|
||
@doc """
|
||
Runs a forward pass with the given parameters and input features.
|
||
|
||
`features` should be an `{batch_size, 13}` tensor of float32.
|
||
Returns an `{batch_size, 1}` tensor of scores in [0, 1].
|
||
"""
|
||
def predict(params, features) do
|
||
model = build()
|
||
{_init_fn, predict_fn} = Axon.build(model)
|
||
predict_fn.(params, features)
|
||
end
|
||
|
||
@doc """
|
||
Returns the ordered list of feature names expected by the model.
|
||
"""
|
||
def feature_names do
|
||
[
|
||
:surface_temp_c,
|
||
:surface_dewpoint_c,
|
||
:surface_pressure_mb,
|
||
:abs_humidity,
|
||
:td_depression,
|
||
:min_refractivity_gradient,
|
||
:hpbl_m,
|
||
:pwat_mm,
|
||
:surface_refractivity,
|
||
:latitude,
|
||
:sfi,
|
||
:kp_max,
|
||
:ducting_detected,
|
||
:k_index,
|
||
:lifted_index,
|
||
:solar_hour_sin,
|
||
:solar_hour_cos,
|
||
:month_sin,
|
||
:month_cos,
|
||
:log_freq_mhz
|
||
]
|
||
end
|
||
|
||
@doc """
|
||
Encodes raw condition data into a feature tensor row.
|
||
|
||
Takes a map of raw values and returns a flat list of 13 floats
|
||
ready to be stacked into a batch tensor.
|
||
"""
|
||
def encode_features(%{} = conditions) do
|
||
temp_c = conditions[:surface_temp_c] || 20.0
|
||
dewpoint_c = conditions[:surface_dewpoint_c] || 10.0
|
||
pressure_mb = conditions[:surface_pressure_mb] || 1013.0
|
||
grad = conditions[:min_refractivity_gradient] || -70.0
|
||
hpbl = conditions[:hpbl_m] || 500.0
|
||
pwat = conditions[:pwat_mm] || 20.0
|
||
refractivity = conditions[:surface_refractivity] || 320.0
|
||
latitude = conditions[:latitude] || 37.0
|
||
sfi = conditions[:sfi] || 120.0
|
||
kp_max = conditions[:kp_max] || 2.0
|
||
ducting = if conditions[:ducting_detected], do: 1.0, else: 0.0
|
||
k_index = conditions[:k_index] || 20.0
|
||
lifted_index = conditions[:lifted_index] || 0.0
|
||
utc_hour = conditions[:utc_hour] || 12
|
||
longitude = conditions[:longitude] || -97.0
|
||
month = conditions[:month] || 6
|
||
freq_mhz = conditions[:freq_mhz] || 10_000
|
||
|
||
abs_humidity = abs_humidity(temp_c, dewpoint_c)
|
||
td_depression = temp_c - dewpoint_c
|
||
|
||
# Cyclical encoding for time features (local solar time)
|
||
local_hour = :math.fmod(utc_hour + longitude / 15 + 24, 24)
|
||
hour_rad = 2 * :math.pi() * local_hour / 24
|
||
month_rad = 2 * :math.pi() * (month - 1) / 12
|
||
|
||
# Normalize all features to ~[0, 1] using physical bounds.
|
||
[
|
||
(temp_c + 40) / 80,
|
||
(dewpoint_c + 40) / 80,
|
||
(pressure_mb - 960) / 80,
|
||
abs_humidity / 25.0,
|
||
td_depression / 40.0,
|
||
(grad + 300) / 300,
|
||
hpbl / 3000.0,
|
||
pwat / 60.0,
|
||
(refractivity - 250) / 120,
|
||
(latitude - 25) / 25,
|
||
(sfi - 60) / 200,
|
||
kp_max / 9.0,
|
||
ducting,
|
||
(k_index + 10) / 50,
|
||
(lifted_index + 10) / 20,
|
||
:math.sin(hour_rad),
|
||
:math.cos(hour_rad),
|
||
:math.sin(month_rad),
|
||
:math.cos(month_rad),
|
||
:math.log(freq_mhz) / 13.0
|
||
]
|
||
end
|
||
|
||
defp abs_humidity(temp_c, dewpoint_c) do
|
||
e_sat = 6.112 * :math.exp(17.67 * dewpoint_c / (dewpoint_c + 243.5))
|
||
217.0 * e_sat / (temp_c + 273.15)
|
||
end
|
||
|
||
@doc """
|
||
Explain a prediction by reporting per-feature sensitivity and
|
||
contribution. Returns a map shaped like:
|
||
|
||
%{
|
||
score: 73,
|
||
contributions: [
|
||
%{feature: :min_refractivity_gradient, normalized: 0.30,
|
||
sensitivity: -22.4, contribution: -4.5},
|
||
...
|
||
]
|
||
}
|
||
|
||
`sensitivity` is the per-feature partial derivative ∂score/∂x̂
|
||
(score points per unit change in the *normalized* feature),
|
||
approximated via a single-batch finite-difference forward pass.
|
||
`contribution` is `sensitivity × (normalized − baseline)`, an
|
||
attribution of the current value to the score relative to a neutral
|
||
reference. Most features baseline at 0.5 (mid-range of the [0,1]
|
||
normalization); the four cyclical sin/cos channels baseline at 0.
|
||
|
||
Contributions are sorted by `|contribution|` descending so the top
|
||
drivers appear first.
|
||
|
||
## Options
|
||
|
||
* `:delta` — perturbation magnitude in normalized units. Defaults
|
||
to 0.05 (5% of the [0,1] feature range), large enough to escape
|
||
float noise on a sigmoid output yet small enough that the
|
||
forward-difference estimate stays close to the local gradient.
|
||
"""
|
||
@spec explain_prediction(function(), Axon.ModelState.t(), map(), keyword()) :: %{
|
||
score: non_neg_integer(),
|
||
contributions: [
|
||
%{
|
||
feature: atom(),
|
||
normalized: float(),
|
||
sensitivity: float(),
|
||
contribution: float()
|
||
}
|
||
]
|
||
}
|
||
def explain_prediction(predict_fn, params, conditions, opts \\ []) do
|
||
delta = Keyword.get(opts, :delta, 0.05)
|
||
|
||
base_features = encode_features(conditions)
|
||
feature_names = feature_names()
|
||
|
||
# Build a {n+1, n} batch: row 0 is the unperturbed input; rows 1..n
|
||
# each bump exactly one feature by `delta`. One forward pass yields
|
||
# the base score plus every per-feature sensitivity.
|
||
perturbed_rows =
|
||
base_features
|
||
|> Enum.with_index()
|
||
|> Enum.map(fn {_, idx} -> List.update_at(base_features, idx, &(&1 + delta)) end)
|
||
|
||
batch = Nx.tensor([base_features | perturbed_rows], type: :f32)
|
||
|
||
outputs =
|
||
params
|
||
|> predict_fn.(%{"features" => batch})
|
||
|> Nx.squeeze(axes: [1])
|
||
|> Nx.to_flat_list()
|
||
|
||
[base_pred | perturbed_preds] = outputs
|
||
|
||
contributions =
|
||
feature_names
|
||
|> Enum.zip(base_features)
|
||
|> Enum.zip(perturbed_preds)
|
||
|> Enum.map(fn {{name, normalized}, perturbed_pred} ->
|
||
# Convert from sigmoid output [0,1] to score points [0,100] and
|
||
# reference each feature against its neutral baseline.
|
||
sensitivity = (perturbed_pred - base_pred) / delta * 100.0
|
||
baseline = baseline_for(name)
|
||
contribution = sensitivity * (normalized - baseline)
|
||
|
||
%{
|
||
feature: name,
|
||
normalized: normalized,
|
||
sensitivity: sensitivity,
|
||
contribution: contribution
|
||
}
|
||
end)
|
||
|> Enum.sort_by(&(-abs(&1.contribution)))
|
||
|
||
score =
|
||
base_pred
|
||
|> Kernel.*(100)
|
||
|> round()
|
||
|> trunc()
|
||
|> max(0)
|
||
|> min(100)
|
||
|
||
%{score: score, contributions: contributions}
|
||
end
|
||
|
||
# Cyclical sin/cos channels are centered at 0 (no rotation); every
|
||
# other normalized feature lives on [0, 1] and a midpoint baseline
|
||
# of 0.5 best represents "neutral" for attribution purposes.
|
||
defp baseline_for(:solar_hour_sin), do: 0.0
|
||
defp baseline_for(:solar_hour_cos), do: 0.0
|
||
defp baseline_for(:month_sin), do: 0.0
|
||
defp baseline_for(:month_cos), do: 0.0
|
||
defp baseline_for(_), do: 0.5
|
||
end
|