feat(ml): retrain propagation model and add per-prediction explainability
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.
This commit is contained in:
parent
f26fbafc29
commit
a6b89961f6
3 changed files with 190 additions and 0 deletions
106
lib_ml/model.ex
106
lib_ml/model.ex
|
|
@ -333,4 +333,110 @@ defmodule Microwaveprop.Propagation.Model do
|
||||||
e_sat = 6.112 * :math.exp(17.67 * dewpoint_c / (dewpoint_c + 243.5))
|
e_sat = 6.112 * :math.exp(17.67 * dewpoint_c / (dewpoint_c + 243.5))
|
||||||
217.0 * e_sat / (temp_c + 273.15)
|
217.0 * e_sat / (temp_c + 273.15)
|
||||||
end
|
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
|
end
|
||||||
|
|
|
||||||
Binary file not shown.
|
|
@ -219,4 +219,88 @@ defmodule Microwaveprop.Propagation.ModelTest do
|
||||||
assert length(Model.feature_names()) == 20
|
assert length(Model.feature_names()) == 20
|
||||||
end
|
end
|
||||||
end
|
end
|
||||||
|
|
||||||
|
describe "explain_prediction/3" do
|
||||||
|
test "returns score and one contribution entry per feature" do
|
||||||
|
params = Model.init()
|
||||||
|
predict_fn = Model.compile_predict()
|
||||||
|
|
||||||
|
result =
|
||||||
|
Model.explain_prediction(predict_fn, params, %{
|
||||||
|
surface_temp_c: 25.0,
|
||||||
|
surface_dewpoint_c: 18.0,
|
||||||
|
surface_pressure_mb: 1013.0,
|
||||||
|
min_refractivity_gradient: -250.0,
|
||||||
|
hpbl_m: 800.0,
|
||||||
|
pwat_mm: 30.0,
|
||||||
|
surface_refractivity: 340.0,
|
||||||
|
latitude: 32.5,
|
||||||
|
longitude: -97.0,
|
||||||
|
sfi: 130.0,
|
||||||
|
kp_max: 2.0,
|
||||||
|
ducting_detected: true,
|
||||||
|
k_index: 25.0,
|
||||||
|
lifted_index: -2.0,
|
||||||
|
utc_hour: 23,
|
||||||
|
month: 7,
|
||||||
|
freq_mhz: 10_000
|
||||||
|
})
|
||||||
|
|
||||||
|
assert is_integer(result.score)
|
||||||
|
assert result.score >= 0 and result.score <= 100
|
||||||
|
assert length(result.contributions) == 20
|
||||||
|
|
||||||
|
Enum.each(result.contributions, fn c ->
|
||||||
|
assert is_atom(c.feature)
|
||||||
|
assert is_float(c.normalized) or is_integer(c.normalized)
|
||||||
|
assert is_float(c.sensitivity)
|
||||||
|
assert is_float(c.contribution)
|
||||||
|
end)
|
||||||
|
end
|
||||||
|
|
||||||
|
test "contributions are sorted by absolute contribution descending" do
|
||||||
|
params = Model.init()
|
||||||
|
predict_fn = Model.compile_predict()
|
||||||
|
|
||||||
|
result =
|
||||||
|
Model.explain_prediction(predict_fn, params, %{
|
||||||
|
surface_temp_c: 15.0,
|
||||||
|
surface_dewpoint_c: 5.0,
|
||||||
|
freq_mhz: 24_000
|
||||||
|
})
|
||||||
|
|
||||||
|
magnitudes = Enum.map(result.contributions, &abs(&1.contribution))
|
||||||
|
assert magnitudes == Enum.sort(magnitudes, :desc)
|
||||||
|
end
|
||||||
|
|
||||||
|
test "every emitted feature name matches feature_names/0" do
|
||||||
|
params = Model.init()
|
||||||
|
predict_fn = Model.compile_predict()
|
||||||
|
|
||||||
|
result = Model.explain_prediction(predict_fn, params, %{freq_mhz: 47_000})
|
||||||
|
|
||||||
|
emitted = result.contributions |> Enum.map(& &1.feature) |> Enum.sort()
|
||||||
|
expected = Enum.sort(Model.feature_names())
|
||||||
|
assert emitted == expected
|
||||||
|
end
|
||||||
|
|
||||||
|
test "perturbing a feature changes the predicted score" do
|
||||||
|
# Verify the sensitivity output is non-trivial: the model must
|
||||||
|
# actually move when at least one feature moves. With random init
|
||||||
|
# this is essentially guaranteed unless every weight collapses.
|
||||||
|
params = Model.init()
|
||||||
|
predict_fn = Model.compile_predict()
|
||||||
|
|
||||||
|
result =
|
||||||
|
Model.explain_prediction(predict_fn, params, %{
|
||||||
|
surface_temp_c: 20.0,
|
||||||
|
surface_dewpoint_c: 10.0,
|
||||||
|
surface_pressure_mb: 1013.0,
|
||||||
|
freq_mhz: 10_000
|
||||||
|
})
|
||||||
|
|
||||||
|
max_sensitivity = result.contributions |> Enum.map(&abs(&1.sensitivity)) |> Enum.max()
|
||||||
|
assert max_sensitivity > 0.0
|
||||||
|
end
|
||||||
|
end
|
||||||
end
|
end
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue