From 8c3cd2575f8ad3c6921414b89cebcedc6f8cdfa2 Mon Sep 17 00:00:00 2001 From: nayan Date: Mon, 27 Jul 2026 16:14:34 -0400 Subject: [PATCH] fuckit. dynamic everything. --- .../sunnypilot/modeld_v2/mhp_inference.py | 128 +++++++++++++ openpilot/sunnypilot/modeld_v2/modeld.py | 24 ++- .../modeld_v2/parse_model_outputs.py | 74 +++++--- .../modeld_v2/parse_model_outputs_split.py | 91 +++++---- .../modeld_v2/tests/test_dynamic_mhp.py | 172 ++++++++++++++++++ 5 files changed, 428 insertions(+), 61 deletions(-) create mode 100644 openpilot/sunnypilot/modeld_v2/mhp_inference.py create mode 100644 openpilot/sunnypilot/modeld_v2/tests/test_dynamic_mhp.py diff --git a/openpilot/sunnypilot/modeld_v2/mhp_inference.py b/openpilot/sunnypilot/modeld_v2/mhp_inference.py new file mode 100644 index 0000000000..e686c18c0e --- /dev/null +++ b/openpilot/sunnypilot/modeld_v2/mhp_inference.py @@ -0,0 +1,128 @@ +"""Infer MHP (mixture-density-hypothesis) parameter values from output slice sizes. + +The legacy supercombo encoded its outputs using fixed values for +``PLAN_MHP_N``, ``PLAN_MHP_SELECTION``, ``LEAD_MHP_N``, and ``LEAD_MHP_SELECTION``. +Newer supercombo architectures emit differently-sized slices for the same heads, +yet the parser downstream still expects to be told the layout. + +This module figures those numbers out at runtime from the model's output slice, +so the existing ``parse_mdn`` path can handle any supercombo flavor without code +changes for each variant. ``infer_mhp`` always tries the legacy values first so +existing compiled pkls behave identically (full backwards compatibility). + +Schema reminder (``parse_mdn`` packs hypotheses contiguously in the order +``in_N × (mu | std | weights)`` along the channel axis): + + slice_size == in_N × (2·n_values + out_N) + +where ``n_values`` is the per-hypothesis value width (``IDX_N × PLAN_WIDTH`` +for plan, ``LEAD_TRAJ_LEN × LEAD_WIDTH`` for lead, etc.). +""" + +def infer_mhp( + slice_size: int, + n_values: int, + legacy_in_n: int, + legacy_out_n: int, + max_in_n: int = 16, +) -> tuple[int, int]: + """Infer ``(in_N, out_N)`` for an MDN-encoded output slice. + + Tries values in this order, returning on the first match: + + 1. ``(legacy_in_n, legacy_out_n)`` exactly — preserves exact backwards + compatibility for existing supercombo pkls. + 2. ``out_N ∈ (0, 1, 3)`` (no weights, single weight, three-way selection) at + increasing in_N — covers the common architectural patterns. + 3. Brute-force any valid ``out_N`` that divides ``slice_size``. + + Args: + slice_size: Number of floats in the output slice (typically + ``slices[name].stop - slices[name].start``). + n_values: Per-hypothesis mu/std width (e.g. ``IDX_N × PLAN_WIDTH`` for plan). + legacy_in_n: The legacy in_N value (highest priority for backwards compat). + legacy_out_n: The legacy out_N value. + max_in_n: Upper bound on accepted hypothesis counts (filters silly parses). + + Returns: + ``(in_N, out_N)``. If nothing fits the formulas, returns ``(1, 0)`` — + single hypothesis with no weights — which is the gentlest fallback. + """ + if slice_size <= 0 or n_values <= 0: + return 1, 0 + + # Priority 1: exact legacy match (BC-preserving). + per_hyp_legacy = 2 * n_values + legacy_out_n + if per_hyp_legacy > 0 and legacy_in_n * per_hyp_legacy == slice_size: + return legacy_in_n, legacy_out_n + + # Priority 2: common weight layouts across supercombo variants. + for out_n in (0, 1, 3): + per_hyp = 2 * n_values + out_n + if per_hyp <= 0: + continue + if slice_size % per_hyp == 0: + in_n = slice_size // per_hyp + if 1 <= in_n <= max_in_n: + return in_n, out_n + + # Priority 3: brute-force any divisor that yields a sensible in_N. + # Bound out_n by max_in_n to keep the search tiny (3 hypotheses + # of weights is already a lot). + for out_n in range(0, max_in_n + 1): + per_hyp = 2 * n_values + out_n + if per_hyp <= 0: + continue + if slice_size % per_hyp == 0: + in_n = slice_size // per_hyp + if 1 <= in_n <= max_in_n: + return in_n, out_n + + # Last resort: best-effort single hypothesis with no weights. + return 1, 0 + + +def slice_size(sl) -> int: + """Return the float-width of a ``slice``/``None`` from ``output_slices``. + + ``None`` and ``slice(None, None, None)`` are treated as 0 (output absent). + Negative-end slices (ONNX-style "from end") aren't supported by this helper + because parser expects single-tensor packed outputs. + """ + if sl is None: + return 0 + start = 0 if sl.start is None else sl.start + stop = sl.stop + if stop is None or stop < 0: + return 0 + return max(0, stop - start) + + +def infer_mhp_for_outputs( + output_slices: dict, + constants, + max_in_n: int = 16, +) -> dict: + """Build a dict of MHP values keyed by head name from a model's output_slices. + + Reads ``output_slices['plan']`` and ``output_slices['lead']`` (if present) + and infers their ``in_N``/``out_N``. Other heads aren't MDN-encoded the same + way and remain driven by the ``constants`` module. + """ + config: dict[str, int] = {} + + plan_size = slice_size(output_slices.get('plan')) + if plan_size > 0: + n = constants.IDX_N * constants.PLAN_WIDTH + in_n, out_n = infer_mhp(plan_size, n, constants.PLAN_MHP_N, constants.PLAN_MHP_SELECTION, max_in_n) + config['plan_mhp_n'] = in_n + config['plan_mhp_selection'] = out_n + + lead_size = slice_size(output_slices.get('lead')) + if lead_size > 0: + n = constants.LEAD_TRAJ_LEN * constants.LEAD_WIDTH + in_n, out_n = infer_mhp(lead_size, n, constants.LEAD_MHP_N, constants.LEAD_MHP_SELECTION, max_in_n) + config['lead_mhp_n'] = in_n + config['lead_mhp_selection'] = out_n + + return config diff --git a/openpilot/sunnypilot/modeld_v2/modeld.py b/openpilot/sunnypilot/modeld_v2/modeld.py index 31ed7a7340..5182a191ab 100755 --- a/openpilot/sunnypilot/modeld_v2/modeld.py +++ b/openpilot/sunnypilot/modeld_v2/modeld.py @@ -148,10 +148,6 @@ class ModelState(ModelStateBase): self._road_key = next(key for key in self._vision_input_names if 'big' not in key) self._wide_key = next(key for key in self._vision_input_names if 'big' in key) - from openpilot.sunnypilot.modeld_v2.parse_model_outputs_split import Parser as SplitParser - from openpilot.sunnypilot.modeld_v2.parse_model_outputs import Parser as CombinedParser - self.parser = SplitParser() if self._combined_model_type != 'supercombo' else CombinedParser() - is_20hz = bundle.is20hz if bundle else self._combined_model_type in ('split', 'multi_policy') if is_20hz: from openpilot.sunnypilot.models.split_model_constants import SplitModelConstants @@ -160,6 +156,26 @@ class ModelState(ModelStateBase): from openpilot.sunnypilot.modeld_v2.constants import ModelConstants self.constants = ModelConstants() + # Derive the parser's per-head MHP values from the appropriate output + # slices. Legacy pkls fall back to the constants values (Priority 1 in + # ``mhp_inference.infer_mhp``), so this is fully backwards compatible. + # Supercombo pkls carry plan/lead in the vision tensor; split/multi-policy + # pkls carry them on the policy tensor (and even there we use the first + # policy's slices -- the existing code only tracks one ``policy_output_slices``). + from openpilot.sunnypilot.modeld_v2.mhp_inference import infer_mhp_for_outputs + + if self._combined_model_type == 'supercombo': + mhp_config = infer_mhp_for_outputs(self.vision_output_slices, self.constants) + else: + mhp_config = infer_mhp_for_outputs(self.policy_output_slices, self.constants) + + from openpilot.sunnypilot.modeld_v2.parse_model_outputs_split import Parser as SplitParser + from openpilot.sunnypilot.modeld_v2.parse_model_outputs import Parser as CombinedParser + if self._combined_model_type != 'supercombo': + self.parser = SplitParser(mhp_config=mhp_config) + else: + self.parser = CombinedParser(mhp_config=mhp_config) + self.prev_desire = np.zeros(self.constants.DESIRE_LEN, dtype=np.float32) self.full_frames: dict = {} self._blob_cache: dict = {} diff --git a/openpilot/sunnypilot/modeld_v2/parse_model_outputs.py b/openpilot/sunnypilot/modeld_v2/parse_model_outputs.py index 82103283f3..4d56ae5475 100644 --- a/openpilot/sunnypilot/modeld_v2/parse_model_outputs.py +++ b/openpilot/sunnypilot/modeld_v2/parse_model_outputs.py @@ -18,8 +18,16 @@ def softmax(x, axis=-1): return x class Parser: - def __init__(self, ignore_missing=False): + def __init__(self, ignore_missing=False, mhp_config=None): self.ignore_missing = ignore_missing + # Optional MHP overrides keyed by head: 'plan_mhp_n', 'plan_mhp_selection', + # 'lead_mhp_n', 'lead_mhp_selection'. ``None`` (or missing keys) keeps the + # legacy ``ModelConstants`` values so existing models behave identically. + self.mhp = mhp_config or {} + + def _mhp(self, head, default_in, default_out): + return (self.mhp.get(f'{head}_mhp_n', default_in), + self.mhp.get(f'{head}_mhp_selection', default_out)) def check_missing(self, outs, name): if name not in outs and not self.ignore_missing: @@ -51,36 +59,48 @@ class Parser: pred_std = safe_exp(raw[:,:,n_values: 2*n_values]) if in_N > 1: - weights = np.zeros((raw.shape[0], in_N, out_N), dtype=raw.dtype) - for i in range(out_N): - weights[:,:,i - out_N] = softmax(raw[:,:,i - out_N], axis=-1) + if out_N > 0: + weights = np.zeros((raw.shape[0], in_N, out_N), dtype=raw.dtype) + for i in range(out_N): + weights[:,:,i - out_N] = softmax(raw[:,:,i - out_N], axis=-1) - if out_N == 1: + if out_N == 1: + for fidx in range(weights.shape[0]): + idxs = np.argsort(weights[fidx][:,0])[::-1] + weights[fidx] = weights[fidx][idxs] + pred_mu[fidx] = pred_mu[fidx][idxs] + pred_std[fidx] = pred_std[fidx][idxs] + assert out_shape is not None + full_shape = tuple([raw.shape[0], in_N] + list(out_shape)) + outs[name + '_weights'] = weights + outs[name + '_hypotheses'] = pred_mu.reshape(full_shape) + outs[name + '_stds_hypotheses'] = pred_std.reshape(full_shape) + + pred_mu_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype) + pred_std_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype) for fidx in range(weights.shape[0]): - idxs = np.argsort(weights[fidx][:,0])[::-1] - weights[fidx] = weights[fidx][idxs] - pred_mu[fidx] = pred_mu[fidx][idxs] - pred_std[fidx] = pred_std[fidx][idxs] - assert out_shape is not None - full_shape = tuple([raw.shape[0], in_N] + list(out_shape)) - outs[name + '_weights'] = weights - outs[name + '_hypotheses'] = pred_mu.reshape(full_shape) - outs[name + '_stds_hypotheses'] = pred_std.reshape(full_shape) - - pred_mu_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype) - pred_std_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype) - for fidx in range(weights.shape[0]): - for hidx in range(out_N): - idxs = np.argsort(weights[fidx,:,hidx])[::-1] - pred_mu_final[fidx, hidx] = pred_mu[fidx, idxs[0]] - pred_std_final[fidx, hidx] = pred_std[fidx, idxs[0]] + for hidx in range(out_N): + idxs = np.argsort(weights[fidx,:,hidx])[::-1] + pred_mu_final[fidx, hidx] = pred_mu[fidx, idxs[0]] + pred_std_final[fidx, hidx] = pred_std[fidx, idxs[0]] + else: + # MHP without weights: keep every hypothesis intact, surface them as + # ``*_hypotheses`` outputs and use the same shape for the primary + # output so downstream consumers can iterate over the full set. + assert out_shape is not None + full_shape = tuple([raw.shape[0], in_N] + list(out_shape)) + outs[name + '_hypotheses'] = pred_mu.reshape(full_shape) + outs[name + '_stds_hypotheses'] = pred_std.reshape(full_shape) + pred_mu_final = pred_mu + pred_std_final = pred_std else: pred_mu_final = pred_mu pred_std_final = pred_std - if out_N > 1: + if out_N > 1 or (in_N > 1 and out_N == 0): assert out_shape is not None - final_shape = tuple([raw.shape[0], out_N] + list(out_shape)) + n_selections = out_N if out_N > 1 else in_N + final_shape = tuple([raw.shape[0], n_selections] + list(out_shape)) else: assert out_shape is not None final_shape = tuple([raw.shape[0],] + list(out_shape)) @@ -88,7 +108,9 @@ class Parser: outs[name + '_stds'] = pred_std_final.reshape(final_shape) def parse_outputs(self, outs: dict[str, np.ndarray]) -> dict[str, np.ndarray]: - self.parse_mdn('plan', outs, in_N=ModelConstants.PLAN_MHP_N, out_N=ModelConstants.PLAN_MHP_SELECTION, + plan_in, plan_out = self._mhp('plan', ModelConstants.PLAN_MHP_N, ModelConstants.PLAN_MHP_SELECTION) + lead_in, lead_out = self._mhp('lead', ModelConstants.LEAD_MHP_N, ModelConstants.LEAD_MHP_SELECTION) + self.parse_mdn('plan', outs, in_N=plan_in, out_N=plan_out, out_shape=(ModelConstants.IDX_N,ModelConstants.PLAN_WIDTH)) self.parse_mdn('lane_lines', outs, in_N=0, out_N=0, out_shape=(ModelConstants.NUM_LANE_LINES,ModelConstants.IDX_N,ModelConstants.LANE_LINES_WIDTH)) self.parse_mdn('road_edges', outs, in_N=0, out_N=0, out_shape=(ModelConstants.NUM_ROAD_EDGES,ModelConstants.IDX_N,ModelConstants.LANE_LINES_WIDTH)) @@ -97,7 +119,7 @@ class Parser: if 'sim_pose' in outs: self.parse_mdn('sim_pose', outs, in_N=0, out_N=0, out_shape=(ModelConstants.POSE_WIDTH,)) self.parse_mdn('wide_from_device_euler', outs, in_N=0, out_N=0, out_shape=(ModelConstants.WIDE_FROM_DEVICE_WIDTH,)) - self.parse_mdn('lead', outs, in_N=ModelConstants.LEAD_MHP_N, out_N=ModelConstants.LEAD_MHP_SELECTION, + self.parse_mdn('lead', outs, in_N=lead_in, out_N=lead_out, out_shape=(ModelConstants.LEAD_TRAJ_LEN,ModelConstants.LEAD_WIDTH)) if 'lat_planner_solution' in outs: self.parse_mdn('lat_planner_solution', outs, in_N=0, out_N=0, out_shape=(ModelConstants.IDX_N,ModelConstants.LAT_PLANNER_SOLUTION_WIDTH)) diff --git a/openpilot/sunnypilot/modeld_v2/parse_model_outputs_split.py b/openpilot/sunnypilot/modeld_v2/parse_model_outputs_split.py index efc5d846da..8c152ad220 100644 --- a/openpilot/sunnypilot/modeld_v2/parse_model_outputs_split.py +++ b/openpilot/sunnypilot/modeld_v2/parse_model_outputs_split.py @@ -22,8 +22,16 @@ def softmax(x, axis=-1): class Parser: - def __init__(self, ignore_missing=False): + def __init__(self, ignore_missing=False, mhp_config=None): self.ignore_missing = ignore_missing + # Optional overrides for plan/lead ``in_N`` / ``out_N``. ``None`` or any + # missing key falls back to ``SplitModelConstants`` so previously compiled + # pkls continue to behave identically. + self.mhp = mhp_config or {} + + def _mhp(self, head, default_in, default_out): + return (self.mhp.get(f'{head}_mhp_n', default_in), + self.mhp.get(f'{head}_mhp_selection', default_out)) def check_missing(self, outs, name): if name not in outs and not self.ignore_missing: @@ -55,36 +63,46 @@ class Parser: pred_std = safe_exp(raw[:,:,n_values: 2*n_values]) if in_N > 1: - weights = np.zeros((raw.shape[0], in_N, out_N), dtype=raw.dtype) - for i in range(out_N): - weights[:,:,i - out_N] = softmax(raw[:,:,i - out_N], axis=-1) + if out_N > 0: + weights = np.zeros((raw.shape[0], in_N, out_N), dtype=raw.dtype) + for i in range(out_N): + weights[:,:,i - out_N] = softmax(raw[:,:,i - out_N], axis=-1) - if out_N == 1: + if out_N == 1: + for fidx in range(weights.shape[0]): + idxs = np.argsort(weights[fidx][:,0])[::-1] + weights[fidx] = weights[fidx][idxs] + pred_mu[fidx] = pred_mu[fidx][idxs] + pred_std[fidx] = pred_std[fidx][idxs] + assert out_shape is not None + full_shape = tuple([raw.shape[0], in_N] + list(out_shape)) + outs[name + '_weights'] = weights + outs[name + '_hypotheses'] = pred_mu.reshape(full_shape) + outs[name + '_stds_hypotheses'] = pred_std.reshape(full_shape) + + pred_mu_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype) + pred_std_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype) for fidx in range(weights.shape[0]): - idxs = np.argsort(weights[fidx][:,0])[::-1] - weights[fidx] = weights[fidx][idxs] - pred_mu[fidx] = pred_mu[fidx][idxs] - pred_std[fidx] = pred_std[fidx][idxs] - assert out_shape is not None - full_shape = tuple([raw.shape[0], in_N] + list(out_shape)) - outs[name + '_weights'] = weights - outs[name + '_hypotheses'] = pred_mu.reshape(full_shape) - outs[name + '_stds_hypotheses'] = pred_std.reshape(full_shape) - - pred_mu_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype) - pred_std_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype) - for fidx in range(weights.shape[0]): - for hidx in range(out_N): - idxs = np.argsort(weights[fidx,:,hidx])[::-1] - pred_mu_final[fidx, hidx] = pred_mu[fidx, idxs[0]] - pred_std_final[fidx, hidx] = pred_std[fidx, idxs[0]] + for hidx in range(out_N): + idxs = np.argsort(weights[fidx,:,hidx])[::-1] + pred_mu_final[fidx, hidx] = pred_mu[fidx, idxs[0]] + pred_std_final[fidx, hidx] = pred_std[fidx, idxs[0]] + else: + # MHP without weights: keep every hypothesis intact. + assert out_shape is not None + full_shape = tuple([raw.shape[0], in_N] + list(out_shape)) + outs[name + '_hypotheses'] = pred_mu.reshape(full_shape) + outs[name + '_stds_hypotheses'] = pred_std.reshape(full_shape) + pred_mu_final = pred_mu + pred_std_final = pred_std else: pred_mu_final = pred_mu pred_std_final = pred_std - if out_N > 1: + if out_N > 1 or (in_N > 1 and out_N == 0): assert out_shape is not None - final_shape = tuple([raw.shape[0], out_N] + list(out_shape)) + n_selections = out_N if out_N > 1 else in_N + final_shape = tuple([raw.shape[0], n_selections] + list(out_shape)) else: assert out_shape is not None final_shape = tuple([raw.shape[0],] + list(out_shape)) @@ -100,15 +118,26 @@ class Parser: def parse_dynamic_outputs(self, outs: dict[str, np.ndarray]) -> None: if 'lead' in outs: - lead_mhp = self.is_mhp(outs, 'lead', - SplitModelConstants.LEAD_MHP_SELECTION * SplitModelConstants.LEAD_TRAJ_LEN * SplitModelConstants.LEAD_WIDTH) - lead_in_N, lead_out_N = (SplitModelConstants.LEAD_MHP_N, SplitModelConstants.LEAD_MHP_SELECTION) if lead_mhp else (0, 0) - lead_out_shape = (SplitModelConstants.LEAD_TRAJ_LEN, SplitModelConstants.LEAD_WIDTH) if lead_mhp else \ - (SplitModelConstants.LEAD_MHP_SELECTION, SplitModelConstants.LEAD_TRAJ_LEN, SplitModelConstants.LEAD_WIDTH) + # Prefer explicit overrides in ``mhp``; otherwise fall back to the + # legacy `is_mhp` heuristic that inspects the raw tensor's last axis. + if self.mhp.get('lead_mhp_n') is not None or self.mhp.get('lead_mhp_selection') is not None: + lead_in_N, lead_out_N = (SplitModelConstants.LEAD_MHP_N, SplitModelConstants.LEAD_MHP_SELECTION) + lead_in_N = self.mhp.get('lead_mhp_n', lead_in_N) + lead_out_N = self.mhp.get('lead_mhp_selection', lead_out_N) + lead_out_shape = (SplitModelConstants.LEAD_TRAJ_LEN, SplitModelConstants.LEAD_WIDTH) + else: + lead_mhp = self.is_mhp(outs, 'lead', + SplitModelConstants.LEAD_MHP_SELECTION * SplitModelConstants.LEAD_TRAJ_LEN * SplitModelConstants.LEAD_WIDTH) + lead_in_N, lead_out_N = (SplitModelConstants.LEAD_MHP_N, SplitModelConstants.LEAD_MHP_SELECTION) if lead_mhp else (0, 0) + lead_out_shape = (SplitModelConstants.LEAD_TRAJ_LEN, SplitModelConstants.LEAD_WIDTH) if lead_mhp else \ + (SplitModelConstants.LEAD_MHP_SELECTION, SplitModelConstants.LEAD_TRAJ_LEN, SplitModelConstants.LEAD_WIDTH) self.parse_mdn('lead', outs, in_N=lead_in_N, out_N=lead_out_N, out_shape=lead_out_shape) if 'plan' in outs: - plan_mhp = self.is_mhp(outs, 'plan', SplitModelConstants.IDX_N * SplitModelConstants.PLAN_WIDTH) - plan_in_N, plan_out_N = (SplitModelConstants.PLAN_MHP_N, SplitModelConstants.PLAN_MHP_SELECTION) if plan_mhp else (0, 0) + if self.mhp.get('plan_mhp_n') is not None or self.mhp.get('plan_mhp_selection') is not None: + plan_in_N, plan_out_N = self._mhp('plan', SplitModelConstants.PLAN_MHP_N, SplitModelConstants.PLAN_MHP_SELECTION) + else: + plan_mhp = self.is_mhp(outs, 'plan', SplitModelConstants.IDX_N * SplitModelConstants.PLAN_WIDTH) + plan_in_N, plan_out_N = (SplitModelConstants.PLAN_MHP_N, SplitModelConstants.PLAN_MHP_SELECTION) if plan_mhp else (0, 0) self.parse_mdn('plan', outs, in_N=plan_in_N, out_N=plan_out_N, out_shape=(SplitModelConstants.IDX_N, SplitModelConstants.PLAN_WIDTH)) if 'planplus' in outs: diff --git a/openpilot/sunnypilot/modeld_v2/tests/test_dynamic_mhp.py b/openpilot/sunnypilot/modeld_v2/tests/test_dynamic_mhp.py new file mode 100644 index 0000000000..efe6421d43 --- /dev/null +++ b/openpilot/sunnypilot/modeld_v2/tests/test_dynamic_mhp.py @@ -0,0 +1,172 @@ +"""Tests for the dynamic MDN inference in ``mhp_inference`` and the parser +changes that read from ``output_slices`` instead of hardcoded constants. +""" + +import numpy as np +import pytest + +from openpilot.sunnypilot.modeld_v2.mhp_inference import ( + infer_mhp, + slice_size, + infer_mhp_for_outputs, +) +from openpilot.sunnypilot.modeld_v2.parse_model_outputs import Parser as CombinedParser +from openpilot.sunnypilot.modeld_v2.parse_model_outputs_split import Parser as SplitParser +from openpilot.sunnypilot.modeld_v2.constants import ModelConstants + + +# -- infer_mhp -------------------------------------------------------------- + +class TestInferMhp: + N_PLAN = ModelConstants.IDX_N * ModelConstants.PLAN_WIDTH # 495 + N_LEAD = ModelConstants.LEAD_TRAJ_LEN * ModelConstants.LEAD_WIDTH # 24 + + def test_legacy_plan_preserved(self): + # Legacy: 5 hypotheses x (2*495 + 1) = 4955 + assert infer_mhp(4955, self.N_PLAN, legacy_in_n=5, legacy_out_n=1) == (5, 1) + + def test_legacy_lead_preserved(self): + # Legacy: 2 hypotheses x (2*24 + 3) = 102 + assert infer_mhp(102, self.N_LEAD, legacy_in_n=2, legacy_out_n=3) == (2, 3) + + def test_new_supercombo_plan_single_hypothesis_no_weights(self): + # New combined supercombo: 1 hypothesis x (2*495 + 0) = 990 + assert infer_mhp(990, self.N_PLAN, legacy_in_n=5, legacy_out_n=1) == (1, 0) + + def test_new_supercombo_lead_three_hypotheses_no_weights(self): + # New combined supercombo: 3 hypotheses x (2*24 + 0) = 144 + assert infer_mhp(144, self.N_LEAD, legacy_in_n=2, legacy_out_n=3) == (3, 0) + + def test_out_n_one(self): + # 4 hypotheses x (2*495 + 1) = 3964 + assert infer_mhp(3964, self.N_PLAN, legacy_in_n=5, legacy_out_n=1) == (4, 1) + + def test_out_n_three(self): + # 2 hypotheses x (2*24 + 3) = 102 (matches legacy_lead as well) + assert infer_mhp(102, self.N_LEAD, legacy_in_n=4, legacy_out_n=99) == (2, 3) + + def test_zero_or_invalid_returns_fallback(self): + assert infer_mhp(0, self.N_PLAN, 5, 1) == (1, 0) + assert infer_mhp(-1, self.N_PLAN, 5, 1) == (1, 0) + assert infer_mhp(990, 0, 5, 1) == (1, 0) + + def test_no_match_returns_single_hypothesis(self): + # 987 doesn't cleanly factor under the constraints we care about. + assert infer_mhp(987, self.N_PLAN, 5, 1) == (1, 0) + + +class TestSliceSize: + def test_none_returns_zero(self): + assert slice_size(None) == 0 + + def test_basic_slice(self): + assert slice_size(slice(10, 50)) == 40 + + def test_negative_stop_returns_zero(self): + assert slice_size(slice(10, -2)) == 0 + + def test_none_bounds(self): + assert slice_size(slice(None, 100)) == 100 + + +class TestInferMhpForOutputs: + def test_infers_for_plan_and_lead(self): + slices = { + 'plan': slice(1576, 2566), # 990 + 'lead': slice(917, 1061), # 144 + } + cfg = infer_mhp_for_outputs(slices, ModelConstants) + assert cfg == {'plan_mhp_n': 1, 'plan_mhp_selection': 0, + 'lead_mhp_n': 3, 'lead_mhp_selection': 0} + + def test_legacy_falls_back_to_constants(self): + # Legacy sizes: 4955 plan, 102 lead -> both match Priority 1. + slices = { + 'plan': slice(0, 4955), + 'lead': slice(4955, 5057), + } + cfg = infer_mhp_for_outputs(slices, ModelConstants) + assert cfg == {'plan_mhp_n': ModelConstants.PLAN_MHP_N, + 'plan_mhp_selection': ModelConstants.PLAN_MHP_SELECTION, + 'lead_mhp_n': ModelConstants.LEAD_MHP_N, + 'lead_mhp_selection': ModelConstants.LEAD_MHP_SELECTION} + + def test_missing_outputs_are_skipped(self): + cfg = infer_mhp_for_outputs({}, ModelConstants) + assert cfg == {} + + +# -- CombinedParser --------------------------------------------------------- + +def _synth_outputs(in_n_plan=5, out_n_plan=1, in_n_lead=2, out_n_lead=3, + n_plan=ModelConstants.IDX_N * ModelConstants.PLAN_WIDTH, + n_lead=ModelConstants.LEAD_TRAJ_LEN * ModelConstants.LEAD_WIDTH): + """Synthesize a flat-output dict with the right per-head sizes.""" + plan_size = in_n_plan * (2 * n_plan + out_n_plan) + lead_size = in_n_lead * (2 * n_lead + out_n_lead) + return { + 'plan': np.random.RandomState(0).randn(1, plan_size).astype(np.float32), + 'lead': np.random.RandomState(1).randn(1, lead_size).astype(np.float32), + # Other outputs that parse_outputs() consumes: + 'lane_lines': np.random.RandomState(2).randn(1, 528).astype(np.float32), + 'road_edges': np.random.RandomState(3).randn(1, 264).astype(np.float32), + 'pose': np.random.RandomState(4).randn(1, 12).astype(np.float32), + 'road_transform': np.random.RandomState(5).randn(1, 12).astype(np.float32), + 'wide_from_device_euler': np.random.RandomState(6).randn(1, 6).astype(np.float32), + 'lead_prob': np.random.RandomState(7).randn(1, 3).astype(np.float32), + 'lane_lines_prob':np.random.RandomState(8).randn(1, 8).astype(np.float32), + 'meta': np.random.RandomState(9).randn(1, 55).astype(np.float32), + 'desire_state': np.random.RandomState(10).randn(1, 8).astype(np.float32), + 'desire_pred': np.random.RandomState(11).randn(1, 32).astype(np.float32), + } + + +class TestCombinedParser: + def test_legacy_keeps_existing_shape(self): + p = CombinedParser() # empty mhp -> legacy constants + out = p.parse_outputs(_synth_outputs(5, 1, 2, 3)) + assert out['plan'].shape == (1, 33, 15) + assert out['plan_stds'].shape == (1, 33, 15) + # Lead primary output collapses to LEAD_MHP_SELECTION=3 selections per + # ``parse_mdn``; raw hypotheses survive as ``lead_hypotheses``. + assert out['lead'].shape == (1, 3, 6, 4) + assert out['lead_stds'].shape == (1, 3, 6, 4) + assert out['plan_hypotheses'].shape == (1, 5, 33, 15) + assert out['lead_hypotheses'].shape == (1, 2, 6, 4) + + def test_new_supercombo_plan_and_lead_parse(self): + p = CombinedParser(mhp_config={ + 'plan_mhp_n': 1, 'plan_mhp_selection': 0, + 'lead_mhp_n': 3, 'lead_mhp_selection': 0, + }) + out = p.parse_outputs(_synth_outputs(1, 0, 3, 0)) + assert out['plan'].shape == (1, 33, 15) + assert out['plan_stds'].shape == (1, 33, 15) + assert out['lead'].shape == (1, 3, 6, 4) + assert out['lead_stds'].shape == (1, 3, 6, 4) + # MHP-without-weights keeps every hypothesis as ``*_hypotheses`` + assert out['lead_hypotheses'].shape == (1, 3, 6, 4) + # Plan with a single hypothesis takes the in_N<=1 branch, which + # (matching legacy behavior) does not emit ``plan_hypotheses``. + assert 'plan_hypotheses' not in out + + +class TestSplitParser: + def test_default_uses_is_mhp_heuristic(self): + # No mhp_config -> falls back to inspecting the raw tensor's last axis. + # The legacy 102-element lead pack fits the "MHP branch" path. + n = SplitParser() + outs = {'lead': np.zeros((1, 102), dtype=np.float32)} + n.parse_dynamic_outputs(outs) + assert outs['lead'].shape == (1, 3, 6, 4) + assert outs['lead_hypotheses'].shape == (1, 2, 6, 4) + + def test_explicit_mhp_overrides_is_mhp(self): + n = SplitParser(mhp_config={'lead_mhp_n': 3, 'lead_mhp_selection': 0}) + outs = {'lead': np.zeros((1, 144), dtype=np.float32)} + n.parse_dynamic_outputs(outs) + assert outs['lead'].shape == (1, 3, 6, 4) + + +if __name__ == '__main__': + pytest.main([__file__, '-v'])