Skip to content

_mti_kernel

tit._mti_kernel

Fused numba kernel for the K>=2 mTI envelope direction search.

This is a per-element re-implementation of the NumPy pipeline in :mod:tit.calc -- _quadratic_forms -> coarse Fibonacci sweep -> _diverse_top_m_directions -> _refine_local_directions -- that walks every element once and keeps all intermediates (the 192-direction sweep, the seed table, the patch evaluations) in small per-element scratch instead of (n, 192) / (n, 6, 16) temporaries. The NumPy path is memory-bound on those temporaries; this kernel is compute-bound and parallel over elements (prange).

Every step mirrors the NumPy code's arithmetic and tie-breaking (first maximum wins in every argmax, the grid-cosine seed exclusion, the patch geometry and the shrinking round schedule), so results agree with the NumPy path to floating-point round-off. :func:tit.calc selects it automatically when numba is importable; see :func:sweep_refine.

The module imports cleanly without numba (HAVE_NUMBA is False and :func:sweep_refine raises), so the caller can fall back.

sweep_refine

sweep_refine(arrs, psi, directions, too_close, patch_weights, n_seeds, refine)

Run the fused sweep(+refine) kernel.

Parameters

arrs : list of np.ndarray, each (N, 3) float64 Validated field list [E_1a, E_1b, ...] (2K arrays). psi : np.ndarray (K,) or None Per-pair envelope phase; None/all-zero selects the real path. directions : np.ndarray (D, 3) Coarse sweep grid (:func:tit.calc._fibonacci_sphere). too_close : np.ndarray (D, D) bool Seed-exclusion table directions @ directions.T > cos(min_angle). patch_weights : np.ndarray (R, patch, 3) Frame weights for every refinement round (already shrunk per round). n_seeds : int refine : bool

Returns

md, carrier_power, best_direction

Source code in tit/_mti_kernel.py
def sweep_refine(arrs, psi, directions, too_close, patch_weights, n_seeds, refine):
    """Run the fused sweep(+refine) kernel.

    Parameters
    ----------
    arrs : list of np.ndarray, each (N, 3) float64
        Validated field list ``[E_1a, E_1b, ...]`` (2K arrays).
    psi : np.ndarray (K,) or None
        Per-pair envelope phase; ``None``/all-zero selects the real path.
    directions : np.ndarray (D, 3)
        Coarse sweep grid (:func:`tit.calc._fibonacci_sphere`).
    too_close : np.ndarray (D, D) bool
        Seed-exclusion table ``directions @ directions.T > cos(min_angle)``.
    patch_weights : np.ndarray (R, patch, 3)
        Frame weights for every refinement round (already shrunk per round).
    n_seeds : int
    refine : bool

    Returns
    -------
    md, carrier_power, best_direction
    """
    if not HAVE_NUMBA:
        raise RuntimeError("numba is not available")
    from tit.calc import _direction_quadratics

    # A homogeneous tuple of C-contiguous (N, 3) float64 arrays: numba
    # indexes it at runtime, so no (N, 2K, 3) stacked copy is needed.
    fields = tuple(np.ascontiguousarray(a, dtype=np.float64) for a in arrs)
    n = fields[0].shape[0]
    n_pairs = len(fields) // 2
    use_phase = bool(psi is not None and np.any(np.asarray(psi) != 0.0))
    if use_phase:
        psi_arr = np.asarray(psi, dtype=np.float64)
        cos_psi = np.cos(psi_arr)
        sin_psi = np.sin(psi_arr)
    else:
        cos_psi = np.ones(n_pairs)
        sin_psi = np.zeros(n_pairs)

    directions = np.ascontiguousarray(directions, dtype=np.float64)
    D6 = np.ascontiguousarray(_direction_quadratics(directions))
    W = np.ascontiguousarray(patch_weights, dtype=np.float64)
    W6 = np.ascontiguousarray(_direction_quadratics(W))
    too_close = np.ascontiguousarray(too_close, dtype=np.bool_)

    md = np.empty(n)
    P = np.empty(n)
    best_dir = np.empty((n, 3))
    _sweep_refine_kernel(
        fields,
        cos_psi,
        sin_psi,
        use_phase,
        directions,
        D6,
        too_close,
        W,
        W6,
        int(n_seeds),
        bool(refine),
        md,
        P,
        best_dir,
    )
    return md, P, best_dir