Solver behind FastRiskScoreClassifier: sparse integer points minimising the calibrated log loss

min over integer w (at most k nonzero, each in [-bound, bound]) of
    min over real a, b of  sum_i log(1 + exp(-y_i (a * x_i.w + b))).

Found by an autoresearch loop (agentic-imodels, evolve_slim) that started from FasterRisk (Liu et al., NeurIPS 2022); this is the final version of run oct05-decile2 (n17_lean).

Pipeline:

  1. data: binary columns are grouped into chains of nested sets (the thresholds of one numeric variable), so each chain is one ordinal code per row; every other column keeps the rank codes of its values. Unique (codes, y) rows with counts; complement pairs of binary columns (x + x' = 1) share one beam hash;
  2. continuous beam search (FasterRisk's): grow the support one column at a time from the best parents. Every column's gain is bounded by one Newton step from per-code histograms; each parent proposes its 10 best columns (at most one per threshold cell: a variable and the decile of the column's row fraction); the best proposals get an exact projected-Newton fit on the parent's row groups;
  3. calibrated rounding: each final support rounded at 20 scales plus 8 scales with the largest points clipped at the bound; the rounding with the smallest calibrated loss is a start;
  4. integer local search from the best 4 starts: value changes and additions, threshold slides (a support threshold moves along its chain), then swaps; moves are ranked by an estimate of the calibrated loss on (score bin, code) cells and the best are checked exactly;
  5. refit swaps: the best swaps at the final point, each followed by a continuous refit of its support and a calibrated rounding, then a local search;
  6. exact polish on near-separable fits (wide logit range and low loss): value changes and screened swaps scored by their exact calibrated loss.

Two settings profiles (PROFILES) differ only in the beam width and diversity rules and a minimum support for binary columns: "decile" for about 9 thresholds per numeric column, "fine" for about 99.

numba is needed to fit but not to import this module: the kernels are plain functions until the first call of solve(), which wraps them with numba.njit and compiles them (or loads them from numba's disk cache) by a warm-up.

Expand source code
"""Solver behind FastRiskScoreClassifier: sparse integer points minimising the calibrated log loss

    min over integer w (at most k nonzero, each in [-bound, bound]) of
        min over real a, b of  sum_i log(1 + exp(-y_i (a * x_i.w + b))).

Found by an autoresearch loop (agentic-imodels, evolve_slim) that started from FasterRisk (Liu et al., NeurIPS
2022); this is the final version of run oct05-decile2 (n17_lean).

Pipeline:

0. data: binary columns are grouped into chains of nested sets (the thresholds of one numeric variable), so each
   chain is one ordinal code per row; every other column keeps the rank codes of its values. Unique (codes, y)
   rows with counts; complement pairs of binary columns (x + x' = 1) share one beam hash;
1. continuous beam search (FasterRisk's): grow the support one column at a time from the best parents. Every
   column's gain is bounded by one Newton step from per-code histograms; each parent proposes its 10 best
   columns (at most one per threshold cell: a variable and the decile of the column's row fraction); the best
   proposals get an exact projected-Newton fit on the parent's row groups;
2. calibrated rounding: each final support rounded at 20 scales plus 8 scales with the largest points clipped
   at the bound; the rounding with the smallest calibrated loss is a start;
3. integer local search from the best 4 starts: value changes and additions, threshold slides (a support
   threshold moves along its chain), then swaps; moves are ranked by an estimate of the calibrated loss on
   (score bin, code) cells and the best are checked exactly;
4. refit swaps: the best swaps at the final point, each followed by a continuous refit of its support and a
   calibrated rounding, then a local search;
5. exact polish on near-separable fits (wide logit range and low loss): value changes and screened swaps scored
   by their exact calibrated loss.

Two settings profiles (PROFILES) differ only in the beam width and diversity rules and a minimum support for
binary columns: "decile" for about 9 thresholds per numeric column, "fine" for about 99.

numba is needed to fit but not to import this module: the kernels are plain functions until the first call of
solve(), which wraps them with numba.njit and compiles them (or loads them from numba's disk cache) by a warm-up.
"""

from __future__ import annotations

import importlib.util
import os
import threading
import time

import numpy as np

#: numba is optional for importing imodels but required to fit this model
HAVE_NUMBA = importlib.util.find_spec("numba") is not None
#: compiled kernels are cached on disk (about 2 minutes to compile once per machine); set
#: RISKSCORE_NUMBA_CACHE=0 to disable, e.g. when the package directory is read-only
NUMBA_CACHE = os.environ.get("RISKSCORE_NUMBA_CACHE", "1") != "0"

# numba and its typed containers, bound by _jit() at the first solve() (the kernels read them as globals)
nb = Dict = List = None
_KERNELS = []  # (name, original Python function, extra njit options) of every kernel, in definition order
# serialises _jit() and _warmup(), so that threads fitting at the same time (GridSearchCV or joblib with the threading
# backend) neither wrap a kernel twice nor run the warm-up twice; reentrant since _warmup() calls _jit()
_LOCK = threading.RLock()


def _kernel(**options):
    """Register a numba kernel. It stays a plain function until _jit() wraps it with numba.njit, so importing this
    module neither imports numba nor sets up one disk cache per kernel."""
    def register(f):
        _KERNELS.append((f.__name__, f, options))
        return f
    return register


def _jit():
    """Replace every registered kernel by its numba dispatcher (once, thread-safe). Kernels call each other through
    these module globals, which numba resolves when it compiles a kernel, i.e. after this. The dispatchers are built
    from the original functions and bound only once all of them exist, so a retry after a failure starts afresh."""
    global nb, Dict, List
    if nb is not None:
        return
    with _LOCK:
        if nb is not None:  # another thread finished while this one waited
            return
        import numba
        from numba.typed import Dict as typed_dict, List as typed_list
        wrapped = {name: numba.njit(cache=NUMBA_CACHE, **options)(f) for name, f, options in _KERNELS}
        globals().update(wrapped)
        Dict, List = typed_dict, typed_list
        nb = numba  # last: other threads read nb to skip the lock

#: settings that differ between the two profiles (everything else is a module constant)
PROFILES = {
    # about 9 thresholds per numeric column (deciles)
    "decile": dict(parent_size=8, final_pool=8, sigmax=99, screen=2.5, lastscreen=3.0, ms_sqrt=0.0),
    # about 99 thresholds per numeric column (percentiles): a wider last level, at most one child per multiset of
    # variables, and binary columns need at least sqrt(n) rows on each side
    "fine": dict(parent_size=10, final_pool=5, sigmax=1, screen=1.5, lastscreen=4.0, ms_sqrt=1.0),
}

COEF_BOUND = 5
TR_A, TR_B = 1.0, 2.0  # trust region of the (a, b) Newton step in the move estimate
DEBRUIJN = np.array([0, 1, 56, 2, 57, 49, 28, 3, 61, 58, 42, 50, 38, 29, 17, 4, 62, 47, 59, 36, 45, 43, 51, 22, 53, 39, 33, 30, 24, 18, 12, 5, 63, 55, 48, 27, 60, 41, 37, 16, 46, 35, 44, 21, 52, 32, 23, 11, 54, 26, 40, 15, 34, 20, 31, 10, 25, 14, 19, 9, 13, 8, 7, 6], np.int64)  # bit index of 2^i by a de Bruijn product
RAW_DENSE = 32  # a non-binary column with at most this many codes is histogrammed, otherwise handled per row
CHILD = 10  # beam: columns proposed per parent
BTOL = 1e-4  # beam: relative tolerance of the projected-Newton child fits
CELLB = 10  # beam: threshold cells per variable (deciles of the row fraction) for the one-per-cell rule
REPR = 0.0911  # beam: children with a point below REPR x the largest on a non-binary column are ranked last
REPP = 1e3  # ... by this loss penalty
NMULT = 20  # rounding scales (largest point 0.5 .. bound + 0.49)
RCLIP, RCLIPN = 4.0, 8  # extra rounding scales with the largest point clipped at the bound, up to RCLIP x
NSTARTS = 4  # local search from at most this many starts ...
SPAT, SPEPS = 2, 0.001  # ... stopping after SPAT starts in a row that do not improve the best by SPEPS (relative)
NEXACT = 2  # local search: candidates checked exactly per move type are 2 x NEXACT
NSCREEN = 3  # swap-in columns screened per removal
ADDSCREEN = 4  # additions screened by a second-order model
SLIDER = 4  # threshold slides: at most this many levels along the chain
SWNG = 1e-9  # swap screen by the continuous Newton gain once the logit range a * (max - min score) reaches this
RSWAP, RSF = 2, 1  # refit swaps: the RSWAP best swaps per round; stop after RSF rounds without improvement
RSEND, RSEN, RSEGAP = 2, 2, 0.003  # refit swaps (RSEN) also from the 2nd best start end within RSEGAP (relative)
POLISH = 3  # exact polish: at most this many rounds ...
POLG = 10.0  # ... when the calibrated map spans a logit range of at least POLG
POLH = 0.35  # ... and the loss is below POLH x the base entropy
POLNS = 16  # polish: swap-in columns per removal state
POLV = 2  # polish: values of each swap-in column scored exactly
POLVV = 3  # polish: values of each support column scored exactly
REJ = 2.0  # exact checks stop once twice the Newton decrement cannot reach the target
CALCAP = 100  # ... or after this many Newton steps
SEPEPS = 1e-7  # stop searching once the calibrated loss is below this (separable)

# Fixed pseudo-random streams (data independent constants): prefixes of default_rng(seed).standard_normal /
# .integers, drawn once at first use, so that a fit does not construct generators.
_RCAP = 1 << 16
_RNORM = {}
_RINT5 = []


def _normals(seed, m):
    if m > _RCAP:
        return np.random.default_rng(seed).standard_normal(m)
    if seed not in _RNORM:
        _RNORM[seed] = np.random.default_rng(seed).standard_normal(_RCAP)
    return _RNORM[seed][:m]


def _ints5(m):
    if m > _RCAP:
        return np.random.default_rng(5).integers(0, 2 ** 63, m, dtype=np.int64)
    if not _RINT5:
        _RINT5.append(np.random.default_rng(5).integers(0, 2 ** 63, _RCAP, dtype=np.int64))
    return _RINT5[0][:m]


# ---------------------------------------------------------------- kernels
@_kernel(inline="always")
def _lrow(m):
    # log(1 + exp(-m)), stable
    if m > 0:
        return np.log1p(np.exp(-m))
    return -m + np.log1p(np.exp(m))


@_kernel(inline="always")
def _sig(z):
    if z >= 0:
        return 1.0 / (1.0 + np.exp(-z))
    e = np.exp(z)
    return e / (1.0 + e)


@_kernel(inline="always")
def _lrow_e(m):
    """(log(1 + exp(-m)), exp(-|m|)): the second gives sigma(-m) without another exp (see _sig_e)."""
    if m > 0:
        e = np.exp(-m)
        return np.log1p(e), e
    e = np.exp(m)
    return -m + np.log1p(e), e


@_kernel(inline="always")
def _sig_e(m, e):
    """sigma(-m) from e = exp(-|m|), bitwise equal to _sig(-m)."""
    if m <= 0:
        return 1.0 / (1.0 + e)
    return e / (1.0 + e)


@_kernel()
def _new_dict_u64():
    """An empty typed dict made inside numba (much cheaper than Dict.empty from Python)."""
    return Dict.empty(key_type=nb.types.uint64, value_type=nb.types.boolean)


@_kernel()
def _new_dict_f64():
    return Dict.empty(key_type=nb.types.float64, value_type=nb.types.boolean)


@_kernel()
def total_loss(ym, c):
    s = 0.0
    for i in range(ym.shape[0]):
        s += c[i] * _lrow(ym[i])
    return s


# ------------------------------------------------------------ data layout
@_kernel()
def complement_rep(B, n, cand, d, R):
    """rep[j]: the smaller index of a complement pair (x_j + x_j' = 1 on every row) of binary columns, else j.
    Columns are matched by a hash of their bitsets and of the complemented bitsets, then checked exactly."""
    nw = B.shape[0]
    rep = np.arange(d)
    last = n - (nw - 1) * 64
    lastmask = np.uint64(0xFFFFFFFFFFFFFFFF) if last == 64 else (np.uint64(1) << np.uint64(last)) - np.uint64(1)
    full = np.uint64(0xFFFFFFFFFFFFFFFF)
    h = Dict.empty(key_type=nb.types.uint64, value_type=nb.types.int64)
    for t in range(cand.shape[0]):
        j = cand[t]
        hv = np.uint64(0)
        for w in range(nw):
            hv += B[w, j] * R[w]
        h[hv] = j
    for t in range(cand.shape[0]):
        j = cand[t]
        hc = np.uint64(0)
        for w in range(nw):
            m = lastmask if w == nw - 1 else full
            hc += ((~B[w, j]) & m) * R[w]
        if hc in h:
            j2 = h[hc]
            ok = True
            for w in range(nw):
                m = lastmask if w == nw - 1 else full
                if ((~B[w, j]) & m) != B[w, j2]:
                    ok = False
                    break
            if ok:
                rep[j] = min(j, j2)
    return rep


@_kernel()
def scan_columns(X):
    """Per column: count of nonzeros, whether a value other than 0 / 1 occurs; bitsets (word, column) of the
    nonzeros (counts by popcount of the bitsets)."""
    n, d = X.shape
    nw = (n + 63) // 64
    B = np.zeros((nw, d), np.uint64)
    bad = np.zeros(d, np.bool_)
    for w in range(nw):
        Bw = B[w]
        for i in range(w * 64, min(n, w * 64 + 64)):
            sh = np.uint64(i & 63)
            row = X[i]
            for j in range(d):
                v = row[j]
                Bw[j] |= np.uint64(v != 0.0) << sh
                bad[j] |= (v != 0.0) & (v != 1.0)
    cnt = np.zeros(d, np.int64)
    for w in range(nw):
        for j in range(d):
            x = B[w, j]
            x = x - ((x >> np.uint64(1)) & np.uint64(0x5555555555555555))
            x = (x & np.uint64(0x3333333333333333)) + ((x >> np.uint64(2)) & np.uint64(0x3333333333333333))
            x = (x + (x >> np.uint64(4))) & np.uint64(0x0F0F0F0F0F0F0F0F)
            cnt[j] += np.int64((x * np.uint64(0x0101010101010101)) >> np.uint64(56))
    return ~bad, cnt, B


@_kernel()
def build_chains(B, cnt, order, d):
    """Greedy chain cover of the columns in `order` (by count descending): each column joins the chain whose
    last column contains it with the smallest count, else starts a new chain. Returns chain id and level
    (1-based position) of every column (-1 when not in `order`), and the number of chains."""
    nw = B.shape[0]
    m = order.shape[0]
    last = np.empty(m, np.int64)
    length = np.zeros(m, np.int64)
    chain = np.full(d, -1, np.int64)
    lev = np.zeros(d, np.int64)
    nch = 0
    for t in range(m):
        j = order[t]
        best = -1
        for ch in range(nch):
            l = last[ch]
            if cnt[l] < cnt[j]:
                continue
            if best >= 0 and cnt[l] >= cnt[last[best]]:
                continue
            ok = True
            for w in range(nw):
                if B[w, j] & ~B[w, l]:
                    ok = False
                    break
            if ok:
                best = ch
        if best < 0:
            best = nch
            nch += 1
        last[best] = j
        length[best] += 1
        chain[j] = best
        lev[j] = length[best]
    return chain, lev, nch


@_kernel()
def chain_codes(B, n, chain, lev, nch, Q):
    """Code of every row in every chain: the largest level whose column contains the row (binary search, the
    columns of a chain are nested)."""
    d = chain.shape[0]
    length = np.zeros(nch, np.int64)
    for j in range(d):
        if chain[j] >= 0:
            length[chain[j]] = max(length[chain[j]], lev[j])
    cp = np.zeros(nch + 1, np.int64)
    for ch in range(nch):
        cp[ch + 1] = cp[ch] + length[ch]
    cols = np.empty(cp[nch], np.int64)
    for j in range(d):
        if chain[j] >= 0:
            cols[cp[chain[j]] + lev[j] - 1] = j
    nw = B.shape[0]
    for ch in range(nch):
        # rows in level l but not in level l + 1 have code l (levels are nested): visit the set bits of the
        # differences, so each row is written once
        m = length[ch]
        base = cp[ch]
        for w in range(nw):
            nxt = np.uint64(0)
            for l in range(m - 1, -1, -1):
                cur = B[w, cols[base + l]]
                x = cur & ~nxt
                nxt = cur
                while x != np.uint64(0):
                    low = x & (~x + np.uint64(1))
                    i = w * 64 + DEBRUIJN[(low * np.uint64(0x03F79D71B4CA8B09)) >> np.uint64(58)]
                    Q[ch, i] = l + 1
                    x ^= low
    return Q


@_kernel()
def colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pow2, vrp, vrows):
    """out[j, p] = sum_i x_ij^(1 + pow2[p]) R[i, p] for every column, from per-code histograms of each variable."""
    V, n = QT.shape
    P = R.shape[1]
    d = lev_of.shape[0]
    out = np.zeros((d, P))
    for v in range(V):
        nc = ncode[v]
        Bq = np.zeros((nc, P))
        for t_ in range(vrp[v], vrp[v + 1]):
            i = vrows[t_]
            q = QT[v, i]
            for p in range(P):
                Bq[q, p] += R[i, p]
        if kind[v] == 0:
            for q in range(nc - 2, -1, -1):  # suffix sums: column of level l sums codes >= l
                for p in range(P):
                    Bq[q, p] += Bq[q + 1, p]
            for t in range(vcp[v], vcp[v + 1]):
                j = vcols[t]
                for p in range(P):
                    out[j, p] = Bq[lev_of[j], p]
        else:
            for t in range(vcp[v], vcp[v + 1]):
                j = vcols[t]
                base = tptr[j]
                for p in range(P):
                    s = 0.0
                    for q in range(nc):
                        x = tabv[base + q]
                        s += (x * x if pow2[p] else x) * Bq[q, p]
                    out[j, p] = s
    return out


@_kernel()
def xval(QT, var_of, tptr, tabv, i, j):
    return tabv[tptr[j] + QT[var_of[j], i]]


@_kernel()
def score_kernel(QT, var_of, tptr, tabv, S, wS):
    n = QT.shape[1]
    s = np.zeros(n)
    for t in range(S.shape[0]):
        j = S[t]
        v = var_of[j]
        base = tptr[j]
        for i in range(n):
            s[i] += wS[t] * tabv[base + QT[v, i]]
    return s


@_kernel()
def chain_tables(tptr, lev_of, var_of, nch):
    """Value tables of the chain columns: x = 1 for codes >= level (other columns filled by the caller)."""
    tabv = np.zeros(tptr[-1])
    for j in range(lev_of.shape[0]):
        if var_of[j] < nch:
            for t in range(tptr[j] + lev_of[j], tptr[j + 1]):
                tabv[t] = 1.0
    return tabv


@_kernel()
def row_hash(Q0, r, ys0):
    V, n = Q0.shape
    h = ys0 * r[V]
    for v in range(V):
        rv = r[v]
        for i in range(n):
            h[i] += rv * Q0[v, i]
    return h


@_kernel()
def unique_rows(h):
    """np.unique(h, return_index=True) with the counts: first occurrence of every distinct value, in value order,
    and the number of rows with that value. A stable LSD radix sort on the order-preserving bit pattern of h
    gives the same order as numpy's stable sort."""
    n = h.shape[0]
    key = np.empty(n, np.uint64)
    hb = h.view(np.uint64)
    top = np.uint64(1) << np.uint64(63)
    for i in range(n):
        u = hb[i]
        key[i] = ~u if (u & top) else (u | top)  # IEEE order -> unsigned order
    o = np.arange(n)
    o2 = np.empty(n, np.int64)
    cnt = np.empty(257, np.int64)
    for sh in range(0, 64, 8):
        cnt[:] = 0
        for t in range(n):
            cnt[((key[o[t]] >> np.uint64(sh)) & np.uint64(255)) + 1] += 1
        if cnt[1:].max() == n:
            continue  # all keys share this byte
        for b in range(256):
            cnt[b + 1] += cnt[b]
        for t in range(n):
            bb = (key[o[t]] >> np.uint64(sh)) & np.uint64(255)
            o2[cnt[bb]] = o[t]
            cnt[bb] += 1
        o, o2 = o2, o
    first = np.empty(n, np.int64)
    cw = np.zeros(n)
    m = -1
    for t in range(n):
        i = o[t]
        if t == 0 or h[i] != h[o[t - 1]]:
            m += 1
            first[m] = i
        cw[m] += 1.0
    m += 1
    return first[:m].copy(), cw[:m].copy()


@_kernel()
def _max_at(a, idx, v):
    """np.maximum.at(a, idx, v) (ufunc.at is slow)."""
    for t in range(idx.shape[0]):
        if v[t] > a[idx[t]]:
            a[idx[t]] = v[t]


@_kernel()
def gather_cols_nz(Q, idx, kind):
    """Q[:, idx] and, per variable, the CSR list of rows that can have a nonzero value (code > 0 for a chain,
    every row otherwise), in two passes (counts while gathering)."""
    V = Q.shape[0]
    m = idx.shape[0]
    out = np.empty((V, m), Q.dtype)
    vrp = np.zeros(V + 1, np.int64)
    for v in range(V):
        cnt = 0
        for t in range(m):
            q = Q[v, idx[t]]
            out[v, t] = q
            cnt += q > 0
        vrp[v + 1] = vrp[v] + (m if kind[v] > 0 else cnt)
    vrows = np.empty(vrp[V], np.int32)
    for v in range(V):
        t0 = vrp[v]
        if kind[v] > 0:
            for i in range(m):
                vrows[t0 + i] = i
        else:
            for i in range(m):
                if out[v, i] > 0:
                    vrows[t0] = i
                    t0 += 1
    return out, vrp, vrows


class Data:
    def colsum(self, R, pow2=False):
        if np.ndim(pow2) == 0:
            pow2 = np.full(R.shape[1], bool(pow2))
        return colsum(self.QT, np.ascontiguousarray(R), self.kind, self.ncode, self.vcp, self.vcols, self.lev_of,
                      self.tptr, self.tabv, np.asarray(pow2, np.bool_), self.vrp, self.vrows)

    def __init__(self, X, y01, ms_sqrt=0.0):
        X = np.ascontiguousarray(X, dtype=np.float64)
        n0, d = X.shape
        self.d = d
        ys0 = np.where(np.asarray(y01) > 0, 1.0, -1.0)
        isbin, cnt, B = scan_columns(X)
        const = np.where(isbin, (cnt == 0) | (cnt == n0), False)
        nb_cols = np.flatnonzero(~isbin)
        if len(nb_cols):
            const[nb_cols] = np.ptp(X[:, nb_cols], axis=0) == 0
        chainable = isbin & ~const
        cand = np.flatnonzero(chainable)
        order = cand[np.lexsort((cand, -cnt[cand]))]
        chain, lev, nch = build_chains(B, cnt, order, d)
        others = np.flatnonzero(~chainable)
        V = nch + len(others)
        kind = np.zeros(V, np.int64)
        ncode = np.zeros(V, np.int64)
        var_of = chain.copy()
        lev_of = lev.copy()
        # codes in a byte when every variable is a chain of at most 255 columns (less memory in every row pass)
        small = len(others) == 0 and (lev.max() if len(lev) else 0) <= 255
        Q0 = np.zeros((V, n0), np.uint8 if small else np.int32)
        chain_codes(B, n0, chain, lev, nch, Q0)
        for ch in range(nch):
            ncode[ch] = 0
        _max_at(ncode, chain[cand], lev[cand])
        ncode[:nch] += 1
        tabs = [None] * d
        for t, j in enumerate(others):
            v = nch + t
            u, inv = np.unique(X[:, j], return_inverse=True)
            Q0[v] = inv.ravel()
            ncode[v] = len(u)
            kind[v] = 1 if len(u) <= RAW_DENSE else 2
            var_of[j] = v
            lev_of[j] = 0
            tabs[j] = u.astype(np.float64)
        tptr = np.zeros(d + 1, np.int64)
        np.cumsum(ncode[var_of], out=tptr[1:])
        tabv = chain_tables(tptr, lev_of, var_of, nch)
        for j in others:
            tabv[tptr[j]:tptr[j + 1]] = tabs[j]
        # unique (codes, y) rows with counts, via a random projection hash (deterministic seed)
        r = _normals(12345, V + 1)
        h = row_hash(Q0, r, ys0)
        first, self.c = unique_rows(h)
        self.QT, vrp_, vrows_ = gather_cols_nz(Q0, first, kind)
        if self.QT.dtype != np.uint8 and ncode.max() <= 256:
            self.QT = self.QT.astype(np.uint8)  # codes in a byte: less memory traffic in every row pass
        self.y = np.ascontiguousarray(ys0[first])
        self.n = self.QT.shape[1]
        self.N = float(self.c.sum())
        self.yc = self.y * self.c
        self.V, self.kind, self.ncode, self.var_of, self.lev_of = V, kind, ncode, var_of, lev_of
        self.tptr, self.tabv = tptr, tabv
        vorder = np.lexsort((lev_of, var_of))
        self.vcols = vorder.astype(np.int64)
        self.vcp = np.zeros(V + 1, np.int64)
        self.vcp[1:] = np.cumsum(np.bincount(var_of, minlength=V))
        self.nchains = nch
        self.allbin = len(others) == 0 or bool(np.all((tabv == 0.0) | (tabv == 1.0)))  # chain tables are 0 / 1
        # rows that can have a nonzero value per variable (code > 0 for a chain, every row otherwise)
        self.vrp, self.vrows = vrp_, vrows_
        cn = self.c / self.N
        if self.allbin:  # x^2 = x for every column: one column sum
            M = self.colsum(cn[:, None], np.array([False]))[:, 0].copy()
            M2 = M
        else:
            MM = self.colsum(np.stack([cn, cn], 1), np.array([False, True]))
            M, M2 = MM[:, 0].copy(), MM[:, 1].copy()
        var = M2 - M * M
        self.norm = np.sqrt(np.maximum(var, 0.0) * self.N)  # centred column norm, as in FasterRisk
        self.valid = self.norm > 1e-9
        # minimum support (profile "fine"): a binary column must hold at least ms_sqrt * sqrt(n) (weighted) rows
        # on each side
        self.cntw = M * self.N
        self.msup = ms_sqrt * np.sqrt(self.N)
        colbin = np.ones(d, np.bool_)
        if not self.allbin:
            colbin[others] = [np.all((tabs[j] == 0) | (tabs[j] == 1)) for j in others]
        self.valid &= ~colbin | ((self.cntw >= self.msup) & (self.N - self.cntw >= self.msup))
        self.bvalid = self.valid.copy()  # columns the beam may add
        # threshold cells for the beam's diversity rules: a chain column's cell is its variable and the decile of
        # its row fraction (CELLB buckets), so thresholds of one variable that split the rows alike share a cell
        frac = self.cntw / self.N
        bucket = np.where(kind[var_of] == 0, np.minimum((frac * CELLB).astype(np.int64), CELLB - 1), 0)
        _, pv_ = np.unique(var_of * CELLB + bucket, return_inverse=True)
        self.pvar = pv_.ravel().astype(np.int64)
        self.npvar = int(self.pvar.max()) + 1 if d else 0
        # complement pairs of binary columns share one beam hash (the same continuous fit up to the intercept)
        self.crep = complement_rep(B, n0, np.flatnonzero(isbin & ~const).astype(np.int64), d,
                                   _ints5(B.shape[0]).astype(np.uint64) * np.uint64(2) + np.uint64(1))


# ------------------------------------------------------------- beam search
@_kernel()
def chol_solve_buf(A, bb, nf, Lm, x):
    """chol_solve on the leading nf x nf block of A, with caller buffers Lm, x (x returned in x[:nf])."""
    for i in range(nf):
        for j in range(i + 1):
            sm = A[i, j]
            for t in range(j):
                sm -= Lm[i, t] * Lm[j, t]
            if i == j:
                if sm <= 0.0:
                    sol = np.linalg.solve(np.ascontiguousarray(A[:nf, :nf]), bb[:nf].copy())
                    for u in range(nf):
                        x[u] = sol[u]
                    return
                Lm[i, i] = np.sqrt(sm)
            else:
                Lm[i, j] = sm / Lm[j, j]
    for i in range(nf):
        x[i] = bb[i]
    for i in range(nf):
        sm = x[i]
        for t in range(i):
            sm -= Lm[i, t] * x[t]
        x[i] = sm / Lm[i, i]
    for i in range(nf - 1, -1, -1):
        sm = x[i]
        for t in range(i + 1, nf):
            sm -= Lm[t, i] * x[t]
        x[i] = sm / Lm[i, i]


@_kernel()
def newton_fit(Z, c, w, lo, hi, maxit, tol):
    """Projected Newton for min sum_g c_g log(1 + exp(-Z_g . w)) with box bounds; w updated in place.
    Returns (loss, margins)."""
    n, p = Z.shape
    m = np.empty(n)
    for i in range(n):
        s = 0.0
        for q in range(p):
            s += Z[i, q] * w[q]
        m[i] = s
    E = np.empty(n)
    En = np.empty(n)
    cur = 0.0
    for i in range(n):
        lv, E[i] = _lrow_e(m[i])
        cur += c[i] * lv
    g = np.empty(p)
    H = np.empty((p, p))
    wn = np.empty(p)
    mn = np.empty(n)
    fi = np.empty(p, np.int64)
    d = np.empty(p)
    A = np.empty((p, p))
    bb = np.empty(p)
    Lm = np.empty((p, p))
    sol = np.empty(p)
    for _ in range(maxit):
        g[:] = 0.0
        H[:, :] = 0.0
        for i in range(n):
            pr = _sig_e(m[i], E[i])
            gi = -c[i] * pr
            hi_ = c[i] * pr * (1.0 - pr)
            for q in range(p):
                zq = Z[i, q]
                g[q] += gi * zq
                hz = hi_ * zq
                for r in range(q, p):
                    H[q, r] += hz * Z[i, r]
        nf = 0
        for q in range(p):
            d[q] = 0.0
            if not ((w[q] <= lo[q] + 1e-12 and g[q] > 0) or (w[q] >= hi[q] - 1e-12 and g[q] < 0)):
                fi[nf] = q
                nf += 1
        if nf > 0:
            for a in range(nf):
                bb[a] = -g[fi[a]]
                for b in range(nf):
                    qa, qb = fi[a], fi[b]
                    A[a, b] = H[min(qa, qb), max(qa, qb)]
                A[a, a] += 1e-10 * (1.0 + A[a, a])
            chol_solve_buf(A, bb, nf, Lm, sol)
            for a in range(nf):
                d[fi[a]] = sol[a]
        t = 1.0
        new = cur
        ok = False
        while t > 1e-8:
            for q in range(p):
                v = w[q] + t * d[q]
                wn[q] = min(max(v, lo[q]), hi[q])
            new = 0.0
            for i in range(n):
                s = 0.0
                for q in range(p):
                    s += Z[i, q] * wn[q]
                mn[i] = s
                lv, En[i] = _lrow_e(s)
                new += c[i] * lv
            if new <= cur:
                ok = True
                break
            t *= 0.5
        if not ok:
            break
        dec = cur - new
        w[:] = wn
        m, mn = mn, m
        E, En = En, E
        cur = new
        if dec <= tol * cur:
            break
    return cur, m


@_kernel()
def regroup(ginv, ng, code, ncode, y, c):
    """Refine row groups by a column's value code. Returns (inv, ng2, group y, group weight, representative row)."""
    n = ginv.shape[0]
    inv = np.empty(n, np.int64)
    M = ng * ncode
    cnt = 0
    if M <= 4 * n + 4096:
        table = np.full(M, -1, np.int64)
        for i in range(n):
            key = ginv[i] * ncode + code[i]
            t = table[key]
            if t < 0:
                t = cnt
                table[key] = t
                cnt += 1
            inv[i] = t
    else:
        keys = np.empty(n, np.int64)
        for i in range(n):
            keys[i] = ginv[i] * ncode + code[i]
        order = np.argsort(keys)
        last = -1
        for t in range(n):
            i = order[t]
            if t == 0 or keys[i] != last:
                cnt += 1
                last = keys[i]
            inv[i] = cnt - 1
    gy = np.empty(cnt)
    gc = np.zeros(cnt)
    rep = np.full(cnt, -1, np.int64)
    for i in range(n):
        g = inv[i]
        gc[g] += c[i]
        if rep[g] < 0:
            rep[g] = i
            gy[g] = y[i]
    return inv, cnt, gy, gc, rep


@_kernel()
def group_design(QT, var_of, tptr, tabv, gy, rep, S):
    ng = rep.shape[0]
    p = S.shape[0] + 1
    Z = np.empty((ng, p))
    for g in range(ng):
        Z[g, 0] = gy[g]
        for q in range(1, p):
            Z[g, q] = gy[g] * xval(QT, var_of, tptr, tabv, rep[g], S[q - 1])
    return Z


@_kernel()
def newton_split(Zp, gy, a0, a1, th, bound, maxit, tol):
    """newton_fit for a child = parent support + one binary column, on the parent's groups: group g has a cell
    with x_j = 0 (weight a0[g]) and one with x_j = 1 (weight a1[g]) that share the parent part of the design, so
    the parent block of the gradient / Hessian and the margins are accumulated once per group. th = (b, parent
    points, new point), updated in place. Returns the loss."""
    ng, ps = Zp.shape
    p = ps + 2
    lo = np.full(p, -bound)
    hi = np.full(p, bound)
    lo[0] = -1e300
    hi[0] = 1e300
    m0 = np.empty(ng)
    E0 = np.empty(ng)
    E1 = np.empty(ng)
    F0 = np.empty(ng)
    F1 = np.empty(ng)
    n0 = np.empty(ng)
    cur = 0.0
    for g in range(ng):
        s = gy[g] * th[0]
        for r in range(ps):
            s += Zp[g, r] * th[r + 1]
        m0[g] = s
        if a0[g] > 0.0:
            lv, E0[g] = _lrow_e(s)
            cur += a0[g] * lv
        if a1[g] > 0.0:
            lv, E1[g] = _lrow_e(s + gy[g] * th[p - 1])
            cur += a1[g] * lv
    gr = np.empty(p)
    H = np.empty((p, p))
    thn = np.empty(p)
    fi = np.empty(p, np.int64)
    d = np.empty(p)
    A = np.empty((p, p))
    bb = np.empty(p)
    Lm = np.empty((p, p))
    sol = np.empty(p)
    zp = np.empty(ps + 1)
    for _ in range(maxit):
        gr[:] = 0.0
        H[:, :] = 0.0
        for g in range(ng):
            yg = gy[g]
            gs = 0.0
            hs = 0.0
            g1 = 0.0
            h1 = 0.0
            if a0[g] > 0.0:
                pr = _sig_e(m0[g], E0[g])
                gs -= a0[g] * pr
                hs += a0[g] * pr * (1.0 - pr)
            if a1[g] > 0.0:
                m1 = m0[g] + yg * th[p - 1]
                pr = _sig_e(m1, E1[g])
                g1 = -a1[g] * pr
                h1 = a1[g] * pr * (1.0 - pr)
                gs += g1
                hs += h1
            zp[0] = yg
            for r in range(ps):
                zp[r + 1] = Zp[g, r]
            for q in range(ps + 1):
                zq = zp[q]
                gr[q] += gs * zq
                hz = hs * zq
                for r in range(q, ps + 1):
                    H[q, r] += hz * zp[r]
                H[q, p - 1] += h1 * zq * yg
            gr[p - 1] += g1 * yg
            H[p - 1, p - 1] += h1 * yg * yg
        nf = 0
        for q in range(p):
            d[q] = 0.0
            if not ((th[q] <= lo[q] + 1e-12 and gr[q] > 0) or (th[q] >= hi[q] - 1e-12 and gr[q] < 0)):
                fi[nf] = q
                nf += 1
        if nf > 0:
            for a in range(nf):
                bb[a] = -gr[fi[a]]
                for b in range(nf):
                    qa, qb = fi[a], fi[b]
                    A[a, b] = H[min(qa, qb), max(qa, qb)]
                A[a, a] += 1e-10 * (1.0 + A[a, a])
            chol_solve_buf(A, bb, nf, Lm, sol)
            for a in range(nf):
                d[fi[a]] = sol[a]
        t = 1.0
        new = cur
        ok = False
        while t > 1e-8:
            for q in range(p):
                v = th[q] + t * d[q]
                thn[q] = min(max(v, lo[q]), hi[q])
            new = 0.0
            for g in range(ng):
                s = gy[g] * thn[0]
                for r in range(ps):
                    s += Zp[g, r] * thn[r + 1]
                n0[g] = s
                if a0[g] > 0.0:
                    lv, F0[g] = _lrow_e(s)
                    new += a0[g] * lv
                if a1[g] > 0.0:
                    lv, F1[g] = _lrow_e(s + gy[g] * thn[p - 1])
                    new += a1[g] * lv
            if new <= cur:
                ok = True
                break
            t *= 0.5
        if not ok:
            break
        dec = cur - new
        th[:] = thn
        m0, n0 = n0, m0
        E0, F0 = F0, E0
        E1, F1 = F1, E1
        cur = new
        if dec <= tol * cur:
            break
    return cur


@_kernel()
def child_fit_batch(par_inv, par_ng, par_gy, par_gc, par_rep, par_S, par_w, js, QT, var_of, lev_of, kind, ncode,
                    tptr, tabv, c, bound, tol, screen, vrp, vrows, init):
    """Fit the children (parent support + column j) for the columns js (sorted by variable) of one parent
    without regrouping rows: the child's cells are the parent's groups split by the value of x_j, with weights
    from a (group, code) histogram of j's variable (suffix sums over codes for a chain)."""
    m = js.shape[0]
    ps = par_S.shape[0]
    losses = np.empty(m)
    Wout = np.empty((m, ps + 2))
    n = par_inv.shape[0]
    Zp = np.empty((par_ng, ps))  # y_g * x_g of the parent's support columns
    for g in range(par_ng):
        for r in range(ps):
            Zp[g, r] = par_gy[g] * xval(QT, var_of, tptr, tabv, par_rep[g], par_S[r])
    offg = np.zeros(par_ng)  # margin of the parent's support part per group (fixed in the screen)
    for g in range(par_ng):
        for r in range(ps):
            offg[g] += Zp[g, r] * par_w[r + 1]
    S_buf = np.empty(ps + 1, np.int64)
    w_buf = np.zeros(ps + 2)
    Z_buf = np.empty((0, 1))
    cw_buf = np.empty(0)
    lo_b = np.full(ps + 2, -bound)
    hi_b = np.full(ps + 2, bound)
    lo_b[0] = -1e300
    hi_b[0] = 1e300
    a0s = np.empty(par_ng)
    a1s = np.empty(par_ng)
    th = np.empty(ps + 2)
    u = 0
    while u < m:
        v = var_of[js[u]]
        u2 = u
        while u2 < m and var_of[js[u2]] == v:
            u2 += 1
        nc = ncode[v]
        H = np.zeros((par_ng, nc))
        if kind[v] == 0 and (u2 - u) * n < n + par_ng * nc:
            # few thresholds of this chain: accumulate the weight of x_j = 1 per group directly
            for t_ in range(vrp[v], vrp[v + 1]):
                i = vrows[t_]
                qi = QT[v, i]
                g = par_inv[i]
                for t in range(u, u2):
                    l = lev_of[js[t]]
                    if qi >= l:
                        H[g, l] += c[i]
        else:
            for t_ in range(vrp[v], vrp[v + 1]):
                i = vrows[t_]
                H[par_inv[i], QT[v, i]] += c[i]
        if kind[v] == 0 and not ((u2 - u) * n < n + par_ng * nc):
            for g in range(par_ng):
                for q in range(nc - 2, -1, -1):
                    H[g, q] += H[g, q + 1]
        for t in range(u, u2):
            j = js[t]
            S = S_buf
            w = w_buf
            w[:] = 0.0
            w[0] = par_w[0]
            pos = 0
            q = 0
            ins = False
            for r in range(ps + 1):
                if not ins and (q >= ps or j < par_S[q]):
                    S[r] = j
                    pos = r
                    w[r + 1] = init[t]
                    ins = True
                else:
                    S[r] = par_S[q]
                    w[r + 1] = par_w[q + 1]
                    q += 1
            if kind[v] == 0 and not screen:
                l = lev_of[j]
                for g in range(par_ng):
                    w1 = H[g, l]
                    a1s[g] = w1
                    w0 = par_gc[g] - w1
                    a0s[g] = w0 if w0 > 0.5 else 0.0
                th[0] = w[0]
                for r in range(ps):
                    th[r + 1] = par_w[r + 1]
                th[ps + 1] = init[t]
                loss = newton_split(Zp, par_gy, a0s, a1s, th, bound, 50, tol)
                w[0] = th[0]
                for r in range(ps + 1):
                    if r < pos:
                        w[r + 1] = th[r + 1]
                    elif r == pos:
                        w[r + 1] = th[ps + 1]
                    else:
                        w[r + 1] = th[r]
                losses[t] = loss
                Wout[t] = w
                continue
            maxcell = par_ng * (2 if kind[v] == 0 else nc)
            ncol = 3 if screen else ps + 2
            if Z_buf.shape[0] < maxcell or Z_buf.shape[1] != ncol:
                Z_buf = np.empty((maxcell, ncol))
                cw_buf = np.empty(maxcell)
            Z = Z_buf
            cw = cw_buf
            nce = 0
            for g in range(par_ng):
                yg = par_gy[g]
                ncg = 2 if kind[v] == 0 else nc
                for xq in range(ncg):
                    if kind[v] == 0:
                        w1 = H[g, lev_of[j]]
                        if xq == 1:
                            wt = w1
                        else:
                            w0 = par_gc[g] - w1
                            wt = w0 if w0 > 0.5 else 0.0
                        xv = float(xq)
                    else:
                        wt = H[g, xq]
                        xv = tabv[tptr[j] + xq]
                    if wt <= 0.0:
                        continue
                    Z[nce, 0] = yg
                    if screen:
                        Z[nce, 1] = yg * xv
                        Z[nce, 2] = offg[g]
                    else:
                        for r in range(ps + 1):
                            if r == pos:
                                Z[nce, r + 1] = yg * xv
                            else:
                                Z[nce, r + 1] = Zp[g, r if r < pos else r - 1]
                    cw[nce] = wt
                    nce += 1
            if screen:
                # only the intercept and the new coefficient move (an upper bound on the child's loss)
                w2 = np.array([w[0], 0.0, 1.0])
                lo2 = np.array([-1e300, -bound, 1.0])
                hi2 = np.array([1e300, bound, 1.0])
                loss, _ = newton_fit(Z[:nce], cw[:nce], w2, lo2, hi2, 20, tol)
                w[0] = w2[0]
                w[pos + 1] = w2[1]
            else:
                loss, _ = newton_fit(Z[:nce], cw[:nce], w, lo_b, hi_b, 50, tol)
            losses[t] = loss
            Wout[t] = w
        u = u2
    return losses, Wout


@_kernel()
def make_child(par_inv, par_ng, par_S, par_w, j, colcode_j, ncode_j, y, c, QT, var_of, tptr, tabv, bound, tol):
    """Add column j to a parent: refine its row groups and refit (b0, beta_S) by projected Newton."""
    inv, ng, gy, gc, rep = regroup(par_inv, par_ng, colcode_j, ncode_j, y, c)
    ps = par_S.shape[0]
    S = np.empty(ps + 1, np.int64)
    w = np.zeros(ps + 2)
    w[0] = par_w[0]
    q = 0
    ins = False
    for r in range(ps + 1):
        if not ins and (q >= ps or j < par_S[q]):
            S[r] = j
            w[r + 1] = 0.0
            ins = True
        else:
            S[r] = par_S[q]
            w[r + 1] = par_w[q + 1]
            q += 1
    lo = np.full(ps + 2, -bound)
    hi = np.full(ps + 2, bound)
    lo[0] = -1e300
    hi[0] = 1e300
    Z = group_design(QT, var_of, tptr, tabv, gy, rep, S)
    loss, mg = newton_fit(Z, gc, w, lo, hi, 50, tol)
    return S, inv, ng, gy, gc, rep, loss, mg, w


@_kernel()
def materialise_kernel(par_inv, par_ng, j, S, w, QT, var_of, kind, lev_of, ncode, tptr, tabv, y, c):
    v = var_of[j]
    n = par_inv.shape[0]
    code = np.empty(n, np.int64)
    if kind[v] == 0:
        # binary column: the parent's groups split by x_j in one pass (groups numbered by first appearance)
        l = lev_of[j]
        table = np.full(2 * par_ng, -1, np.int64)
        inv = np.empty(n, np.int64)
        gy = np.empty(n)
        gc = np.zeros(n)
        rep = np.empty(n, np.int64)
        ng = 0
        for i in range(n):
            key = 2 * par_inv[i] + (1 if QT[v, i] >= l else 0)
            g = table[key]
            if g < 0:
                g = ng
                table[key] = g
                rep[g] = i
                gy[g] = y[i]
                ng += 1
            inv[i] = g
            gc[g] += c[i]
        gy = gy[:ng].copy()
        gc = gc[:ng].copy()
        rep = rep[:ng].copy()
        Z = group_design(QT, var_of, tptr, tabv, gy, rep, S)
        mg = np.zeros(ng)
        for g in range(ng):
            for q in range(Z.shape[1]):
                mg[g] += Z[g, q] * w[q]
        return inv, ng, gy, gc, rep, mg
    else:
        for i in range(n):
            code[i] = QT[v, i]
        nc = ncode[v]
    inv, ng, gy, gc, rep = regroup(par_inv, par_ng, code, nc, y, c)
    Z = group_design(QT, var_of, tptr, tabv, gy, rep, S)
    mg = np.zeros(ng)
    for g in range(ng):
        for q in range(Z.shape[1]):
            mg[g] += Z[g, q] * w[q]
    return inv, ng, gy, gc, rep, mg


@_kernel()
def pick_per_var(g, var_of, V, m):
    """Columns with the largest g > 0, at most one per variable (its best), best first."""
    bestj = np.full(V, -1, np.int64)
    for j in range(g.shape[0]):
        if g[j] > 0:
            v = var_of[j]
            if bestj[v] < 0 or g[j] > g[bestj[v]]:
                bestj[v] = j
    nv = 0
    for v in range(V):
        if bestj[v] >= 0:
            nv += 1
    cand = np.empty(nv, np.int64)
    vals = np.empty(nv)
    u = 0
    for v in range(V):
        if bestj[v] >= 0:
            cand[u] = bestj[v]
            vals[u] = -g[bestj[v]]
            u += 1
    o = np.argsort(vals, kind="mergesort")
    return cand[o[:m]]


@_kernel()
def beam_rows_pr2(Pr, yc, c, third):
    """beam_rows from sigma(-m) per row; the x^2 block only when `third`; also the column sums of the curvature."""
    n, P = Pr.shape
    R = np.empty((n, (3 if third else 2) * P))
    h00 = np.zeros(P)
    for i in range(n):
        for q in range(P):
            pr = Pr[i, q]
            R[i, q] = yc[i] * pr
            h = c[i] * pr * (1.0 - pr)
            R[i, P + q] = h
            h00[q] += h
            if third:
                R[i, 2 * P + q] = h
    return R, h00


@_kernel()
def newton_gains(Gm, h00, valid, bound):
    """gain[q, j]: loss decrease of one Newton step on a new coefficient for column j (box-clipped), with the
    intercept's curvature projected out."""
    d = Gm.shape[0]
    P = h00.shape[0]
    GA = np.full((P, d), -1.0)
    BA = np.zeros((P, d))
    for j in range(d):
        if not valid[j]:
            continue
        for q in range(P):
            g = abs(Gm[j, q])
            h0 = Gm[j, P + q]
            heff = max(Gm[j, 2 * P + q] - h0 * h0 / max(h00[q], 1e-300), 1e-12)
            if g <= bound * heff:
                GA[q, j] = 0.5 * g * g / heff
                BA[q, j] = Gm[j, q] / heff
            else:
                GA[q, j] = bound * g - 0.5 * heff * bound * bound
                BA[q, j] = bound if Gm[j, q] > 0 else -bound
    return GA, BA


@_kernel()
def beam_proposals(GA, BA, ploss, PS, plen, phash, psig, colh, varh, var_of, V, kind, vcols, vcp, seen,
                   child_size, nfit, sigmax):
    """Per parent its child_size best variables (each at its best column by the Newton gain), new supports only
    (seen: hashes of supports already proposed); sorted by estimated loss, at most sigmax per multiset of
    variables, the nfit best. Returns (parent, column, estimated loss, warm start, support hash, signature)."""
    P, d = GA.shape
    maxp = P * child_size
    e_est = np.empty(maxp)
    e_q = np.empty(maxp, np.int64)
    e_j = np.empty(maxp, np.int64)
    e_b = np.empty(maxp)
    e_h = np.empty(maxp, np.uint64)
    e_s = np.empty(maxp, np.uint64)
    u = 0
    for q in range(P):
        gain = GA[q].copy()
        Sq = PS[q, :plen[q]]
        for t in range(plen[q]):
            gain[Sq[t]] = -1.0
        pick = pick_per_var(gain, var_of, V, child_size)  # at most one column per threshold cell
        for t in range(pick.shape[0]):
            j = pick[t]
            key = phash[q] + colh[j]
            if key in seen:
                continue
            seen[key] = True
            e_est[u] = ploss[q] - gain[j]
            e_q[u] = q
            e_j[u] = j
            e_b[u] = BA[q, j]
            e_h[u] = key
            e_s[u] = psig[q] + varh[var_of[j]]
            u += 1
    o = np.argsort(e_est[:u], kind="mergesort")
    cnt = Dict.empty(key_type=nb.types.uint64, value_type=nb.types.int64)
    sel = np.empty(min(u, nfit), np.int64)
    m = 0
    for t in range(u):
        if m >= nfit:
            break
        i = o[t]
        c0 = cnt.get(e_s[i], 0)
        if c0 < sigmax:
            cnt[e_s[i]] = c0 + 1
            sel[m] = i
            m += 1
    sel = sel[:m]
    return e_q[sel], e_j[sel], e_est[sel], e_b[sel], e_h[sel], e_s[sel]


@_kernel()
def beam_level(PINV, PNG, GOFF, GY, GC, GREP, GMG, PS, plen, PW, ploss, phash, psig, QT, var_of, lev_of, kind,
               ncode, vcp, vcols, tptr, tabv, vrp, vrows, bvalid, y, c, yc, allbin, colh, varh, pvar, seen, V,
               child_size, nfit, keep, sigmax, bound, tol):
    """One beam level on parents stored as arrays (row groups concatenated, offsets GOFF): Newton-gain screen of
    every column for every parent, proposals, exact child fits, sort by loss, at most sigmax children per
    signature, the keep best materialised as the next parents. Returns ok = False when nothing is proposed."""
    P = PNG.shape[0]
    n = PINV.shape[1]
    d = var_of.shape[0]
    Pr = np.empty((n, P))
    for q in range(P):
        off = GOFF[q]
        sg = np.empty(PNG[q])
        for g in range(PNG[q]):
            sg[g] = _sig(-GMG[off + g])
        for i in range(n):
            Pr[i, q] = sg[PINV[q, i]]
    R, h00 = beam_rows_pr2(Pr, yc, c, not allbin)
    if allbin:  # x^2 = x: the third block equals the second
        G2 = colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, np.zeros(2 * P, np.bool_), vrp, vrows)
        Gm = np.empty((d, 3 * P))
        Gm[:, :2 * P] = G2
        Gm[:, 2 * P:] = G2[:, P:]
    else:
        pw = np.zeros(3 * P, np.bool_)
        pw[2 * P:] = True
        Gm = colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pw, vrp, vrows)
    GA, BA = newton_gains(Gm, h00, bvalid, bound)
    eq, ej, eest, eb, eh, es = beam_proposals(GA, BA, ploss, PS, plen, phash, psig, colh, varh, pvar, V, kind,
                                              vcols, vcp, seen, child_size, nfit, sigmax)
    m = eq.shape[0]
    lmax = PS.shape[1]
    if m == 0:
        return (False, PINV, PNG, GOFF, GY, GC, GREP, GMG, PS, plen, PW, ploss, phash, psig)
    # exact child fits, in the order: parents ascending; per parent its binary / small-code columns sorted by
    # variable (stable), then the others
    closs = np.empty(m)
    cq = np.empty(m, np.int64)
    cj = np.empty(m, np.int64)
    ch_ = np.empty(m, np.uint64)
    cs_ = np.empty(m, np.uint64)
    CW = np.zeros((m, lmax + 2))
    CS = np.zeros((m, lmax + 1), np.int64)
    u = 0
    for q in range(P):
        cnt = 0
        for t in range(m):
            if eq[t] == q:
                cnt += 1
        if cnt == 0:
            continue
        lst = np.empty(cnt, np.int64)
        cnt = 0
        for t in range(m):
            if eq[t] == q:
                lst[cnt] = t
                cnt += 1
        off = GOFF[q]
        ng = PNG[q]
        par_S = PS[q, :plen[q]].copy()
        par_w = PW[q, :plen[q] + 1].copy()
        gy = GY[off:off + ng]
        gc = GC[off:off + ng]
        rep = GREP[off:off + ng]
        nsp = 0
        for t in range(cnt):
            if kind[var_of[ej[lst[t]]]] < 2:
                nsp += 1
        if nsp > 0:
            idx = np.empty(nsp, np.int64)
            vv = np.empty(nsp, np.int64)
            nsp = 0
            for t in range(cnt):
                if kind[var_of[ej[lst[t]]]] < 2:
                    idx[nsp] = lst[t]
                    vv[nsp] = var_of[ej[lst[t]]]
                    nsp += 1
            idx = idx[np.argsort(vv, kind="mergesort")]
            js = ej[idx].copy()
            b0s = eb[idx].copy()
            losses, W = child_fit_batch(PINV[q], ng, gy, gc, rep, par_S, par_w, js, QT, var_of, lev_of, kind, ncode,
                                        tptr, tabv, c, bound, tol, False, vrp, vrows, b0s)
            for t in range(nsp):
                closs[u] = losses[t]
                cq[u] = q
                cj[u] = js[t]
                ch_[u] = eh[idx[t]]
                cs_[u] = es[idx[t]]
                CW[u, :plen[q] + 2] = W[t]
                u += 1
        for t in range(cnt):
            tt = lst[t]
            j = ej[tt]
            if kind[var_of[j]] < 2:
                continue
            v = var_of[j]
            code = np.empty(n, np.int64)
            for i in range(n):
                code[i] = QT[v, i]
            S, inv_, ng_, gy_, gc_, rep_, loss, mg_, w = make_child(PINV[q], ng, par_S, par_w, j, code, ncode[v], y,
                                                                    c, QT, var_of, tptr, tabv, bound, tol)
            closs[u] = loss
            cq[u] = q
            cj[u] = j
            ch_[u] = eh[tt]
            cs_[u] = es[tt]
            CW[u, :plen[q] + 2] = w
            u += 1
    # supports of the children (sorted)
    for t in range(m):
        q = cq[t]
        ps = plen[q]
        j = cj[t]
        r = 0
        ins = False
        for z in range(ps + 1):
            if not ins and (r >= ps or j < PS[q, r]):
                CS[t, z] = j
                ins = True
            else:
                CS[t, z] = PS[q, r]
                r += 1
    # a child with a coefficient that rounds to 0 at every scale of the rounding grid is not representable by
    # integer points (its support collapses to a smaller one after rounding): rank it after the others, when that
    # column is non-binary (e.g. a raw column on a large scale)
    for t in range(m):
        ps = plen[cq[t]] + 1
        top = 0.0
        for r in range(1, ps + 1):
            top = max(top, abs(CW[t, r]))
        for r in range(1, ps + 1):
            if abs(CW[t, r]) < REPR * top and kind[var_of[CS[t, r - 1]]] != 0:
                closs[t] += REPP
                break
    o = np.argsort(closs, kind="mergesort")
    sel = np.empty(min(m, keep), np.int64)
    ns = 0
    if sigmax > 0:
        sc = Dict.empty(key_type=nb.types.uint64, value_type=nb.types.int64)
        for t in range(m):
            if ns >= keep:
                break
            i = o[t]
            c0 = sc.get(cs_[i], 0)
            if c0 < sigmax:
                sc[cs_[i]] = c0 + 1
                sel[ns] = i
                ns += 1
    else:
        for t in range(min(m, keep)):
            sel[t] = o[t]
        ns = min(m, keep)
    sel = sel[:ns]
    # materialise the kept children as the next parents
    P2 = ns
    PINV2 = np.empty((P2, n), np.int64)
    PNG2 = np.empty(P2, np.int64)
    GOFF2 = np.zeros(P2 + 1, np.int64)
    PS2 = np.zeros((P2, lmax + 1), np.int64)
    plen2 = np.empty(P2, np.int64)
    PW2 = np.zeros((P2, lmax + 2))
    ploss2 = np.empty(P2)
    phash2 = np.empty(P2, np.uint64)
    psig2 = np.empty(P2, np.uint64)
    parts = List()
    for t in range(P2):
        i = sel[t]
        q = cq[i]
        ps = plen[q] + 1
        S = CS[i, :ps].copy()
        w = CW[i, :ps + 1].copy()
        inv_, ng_, gy_, gc_, rep_, mg_ = materialise_kernel(PINV[q], PNG[q], cj[i], S, w, QT, var_of, kind, lev_of,
                                                            ncode, tptr, tabv, y, c)
        PINV2[t] = inv_
        PNG2[t] = ng_
        GOFF2[t + 1] = GOFF2[t] + ng_
        PS2[t, :ps] = S
        plen2[t] = ps
        PW2[t, :ps + 1] = w
        ploss2[t] = closs[i]
        phash2[t] = ch_[i]
        psig2[t] = cs_[i]
        parts.append((gy_, gc_, rep_, mg_))
    tot = GOFF2[P2]
    GY2 = np.empty(tot)
    GC2 = np.empty(tot)
    GREP2 = np.empty(tot, np.int64)
    GMG2 = np.empty(tot)
    for t in range(P2):
        gy_, gc_, rep_, mg_ = parts[t]
        a0 = GOFF2[t]
        GY2[a0:a0 + PNG2[t]] = gy_
        GC2[a0:a0 + PNG2[t]] = gc_
        GREP2[a0:a0 + PNG2[t]] = rep_
        GMG2[a0:a0 + PNG2[t]] = mg_
    return (True, PINV2, PNG2, GOFF2, GY2, GC2, GREP2, GMG2, PS2, plen2, PW2, ploss2, phash2, psig2)


@_kernel()
def beam_levels(st, lev0, nlev, QT, var_of, lev_of, kind, ncode, vcp, vcols, tptr, tabv, vrp, vrows, bvalid, y,
                c, yc, allbin, colh, varh, pvar, seen, V, child_size, parent_size, final_pool, screen, lastscreen, sigmax,
                bound, tol):
    """Beam levels lev0 .. nlev - 1 in one kernel (same steps as the per-level loop of beam_search)."""
    for lev in range(lev0, nlev):
        last = lev == nlev - 1
        keep = final_pool if last else parent_size
        nfit = int((lastscreen if last else screen) * keep)
        w_ = max(lev, 1)
        res = beam_level(st[0], st[1], st[2], st[3], st[4], st[5], st[6], np.ascontiguousarray(st[7][:, :w_]),
                         st[8], np.ascontiguousarray(st[9][:, :w_ + 1]), st[10], st[11], st[12], QT, var_of, lev_of,
                         kind, ncode, vcp, vcols, tptr, tabv, vrp, vrows, bvalid, y, c, yc, allbin, colh, varh,
                         pvar, seen, V, child_size, nfit, keep, sigmax, bound, tol)
        if not res[0]:
            break
        st = (res[1], res[2], res[3], res[4], res[5], res[6], res[7], res[8], res[9], res[10], res[11], res[12],
              res[13])
    return st


def beam_search(D, k, bound, parent_size, final_pool, screen, lastscreen, sigmax, child_size=CHILD,
                deadline=np.inf):
    """Beam search over supports, one numba kernel per level (beam_level). Every column's gain for every parent is
    bounded by one Newton step on the new coefficient (from per-code histograms of the gradient and curvature
    rows); each parent proposes its child_size best variables (each at its best column), and only the best
    screen x parent_size proposals (lastscreen x final_pool at the last level) are fitted exactly; at most sigmax
    children per multiset of variables are kept. Returns the final beam and whether it ran out of time."""
    pvar, PV = D.pvar, D.npvar  # the variable (or the threshold cell of a variable) each column belongs to
    hi5 = _ints5(D.d + max(D.V, PV))
    colh = hi5[:D.d].astype(np.uint64) * np.uint64(2) + np.uint64(1)
    colh = colh[D.crep]  # a complement pair has one hash: mirrored supports are proposed once
    varh = hi5[D.d:D.d + PV].astype(np.uint64) * np.uint64(2) + np.uint64(1)
    inv, ng, gy, gc, rep = regroup(np.zeros(D.n, np.int64), 1, (D.y > 0).astype(np.int64), 2, D.y, D.c)
    npos = D.c[D.y > 0].sum()
    w0 = np.log(npos / (D.N - npos))
    mg = gy * w0
    lmax = max(k, 1)
    st = (inv[None, :].copy(), np.array([ng], np.int64), np.array([0, ng], np.int64), gy, gc, rep, mg,
          np.zeros((1, lmax), np.int64), np.zeros(1, np.int64), np.full((1, lmax + 1), w0),
          np.array([total_loss(mg, gc)]), np.zeros(1, np.uint64), np.zeros(1, np.uint64))
    seen = _new_dict_u64()
    nlev = min(k, int(D.valid.sum()))
    t_start = time.perf_counter()
    stopped = False
    for lev in range(nlev):
        if lev == 1:
            # the remaining levels in one kernel when they surely fit in the time left (level 0, one parent, took
            # dt; a level of the full beam costs at most ~parent_size times that)
            dt = time.perf_counter() - t_start
            if t_start + 2.0 * parent_size * nlev * max(dt, 1e-3) < deadline:
                return False, beam_levels(st, 1, nlev, D.QT, D.var_of, D.lev_of, D.kind, D.ncode, D.vcp, D.vcols,
                                          D.tptr, D.tabv, D.vrp, D.vrows, D.bvalid, D.y, D.c, D.yc, D.allbin, colh,
                                          varh, pvar, seen, PV, child_size, parent_size, final_pool, screen,
                                          lastscreen, sigmax, bound, BTOL)
        last = lev == nlev - 1
        if time.perf_counter() > deadline and st[1].shape[0] > 1:
            # out of time: finish the support greedily from the best parent
            stopped = True
            ng0 = st[1][0]
            st = (st[0][:1], st[1][:1], st[2][:2], st[3][:ng0], st[4][:ng0], st[5][:ng0], st[6][:ng0], st[7][:1],
                  st[8][:1], st[9][:1], st[10][:1], st[11][:1], st[12][:1])
        keep = final_pool if last else parent_size
        nfit = int((lastscreen if last else screen) * keep)
        res = beam_level(st[0], st[1], st[2], st[3], st[4], st[5], st[6], np.ascontiguousarray(st[7][:, :max(lev, 1)]),
                         st[8], np.ascontiguousarray(st[9][:, :max(lev, 1) + 1]), st[10], st[11], st[12], D.QT,
                         D.var_of, D.lev_of, D.kind, D.ncode, D.vcp, D.vcols, D.tptr, D.tabv, D.vrp, D.vrows,
                         D.bvalid, D.y, D.c, D.yc, D.allbin, colh, varh, pvar, seen, PV, child_size, nfit, keep, sigmax,
                         bound, BTOL)
        if not res[0]:
            break
        st = res[1:]
    return stopped, st


# ------------------------------------------------------- calibrated rounding
@_kernel(inline="always")
def _gcd(a, b):
    while b:
        a, b = b, a % b
    return a


@_kernel()
def calib_round_kernel(XS, gy, gc, beta, n_mult, bound):
    """Round m * beta for a grid of scales (largest point 0.5 .. bound + 0.49); return the rounding with the
    smallest calibrated loss on the row groups (XS: group rows of the support columns)."""
    ng, p = XS.shape
    top = 0.0
    for q in range(p):
        top = max(top, abs(beta[q]))
    best_l = np.inf
    best_r = np.zeros(p)
    prev = np.full(p, np.nan)
    r = np.empty(p)
    sg = np.empty(ng)
    if top < 1e-12:
        return best_r, best_l
    seen_r = np.empty((n_mult + RCLIPN, p))
    nseen = 0
    WPc = np.zeros(0)
    WNc = np.zeros(0)
    pres = np.zeros(0, np.bool_)
    svb = np.empty(0)
    Wpb = np.empty(0)
    Wnb = np.empty(0)
    Eb = np.empty(0)
    Enb = np.empty(0)
    for t in range(n_mult + RCLIPN):
        if t < n_mult:
            L = 0.5 + (bound - 0.01) * t / max(n_mult - 1, 1)
        else:
            # clipped scales: the largest points saturate at the bound, the others get finer ratios
            L = bound * (1.0 + (RCLIP - 1.0) * (t - n_mult + 1) / RCLIPN)
        same = True
        anynz = False
        for q in range(p):
            v = np.round(beta[q] * L / top)
            v = min(max(v, -bound), bound)
            r[q] = v
            if v != prev[q]:
                same = False
            if v != 0:
                anynz = True
        if same or not anynz:
            continue
        prev[:] = r
        # a rounding proportional to an earlier one has the same calibrated loss: skip it
        gg = 0
        for q in range(p):
            gg = _gcd(gg, int(abs(r[q])))
        sgn = 0.0
        for q in range(p):
            if r[q] != 0.0:
                sgn = 1.0 if r[q] > 0 else -1.0
                break
        dup = False
        for t2 in range(nseen):
            eq_ = True
            for q in range(p):
                if seen_r[t2, q] != sgn * r[q] / gg:
                    eq_ = False
                    break
            if eq_:
                dup = True
                break
        if dup:
            continue
        for q in range(p):
            seen_r[nseen, q] = sgn * r[q] / gg
        nseen += 1
        mn = np.inf
        mx = -np.inf
        m1 = 0.0
        m2 = 0.0
        for g in range(ng):
            v = 0.0
            for q in range(p):
                v += XS[g, q] * r[q]
            sg[g] = v
            mn = min(mn, v)
            mx = max(mx, v)
            m1 += gc[g] * v
            m2 += gc[g] * v * v
        if mx == mn:
            continue
        tot = gc.sum()
        sd = np.sqrt(max(m2 / tot - (m1 / tot) ** 2, 1e-300))
        # integer scores: calibrate on the distinct score values (counted in reused buffers)
        integral = True
        for g in range(ng):
            if sg[g] != np.floor(sg[g]):
                integral = False
                break
        if integral and mx - mn <= 4 * ng + 1024:
            R = int(mx - mn) + 1
            if R > WPc.shape[0]:
                WPc = np.zeros(R)
                WNc = np.zeros(R)
                pres = np.zeros(R, np.bool_)
                svb = np.empty(R)
                Wpb = np.empty(R)
                Wnb = np.empty(R)
                Eb = np.empty(2 * R)
                Enb = np.empty(2 * R)
            for t2 in range(R):
                WPc[t2] = 0.0
                WNc[t2] = 0.0
                pres[t2] = False
            for g in range(ng):
                ix = int(sg[g] - mn)
                pres[ix] = True
                if gy[g] > 0:
                    WPc[ix] += gc[g]
                else:
                    WNc[ix] += gc[g]
            mb_ = 0
            for t2 in range(R):
                if pres[t2]:
                    svb[mb_] = (mn + t2) / sd
                    Wpb[mb_] = WPc[t2]
                    Wnb[mb_] = WNc[t2]
                    mb_ += 1
            loss, _, _ = calibrate_bins_buf(svb[:mb_], Wpb[:mb_], Wnb[:mb_], 0.0, 0.0, Eb, Enb)
        else:
            inv_, sv, Wp, Wn = bin_scores(sg, gy, gc)
            loss, _, _ = calibrate_bins(sv / sd, Wp, Wn, 0.0, 0.0)
        if loss < best_l:
            best_l = loss
            best_r[:] = r
    return best_r, best_l


@_kernel()
def round_all(GOFF, PNG, GY, GC, GREP, PS, plen, PW, QT, var_of, tptr, tabv, n_mult, bound, hv, N):
    """calib_round for every final beam node: rounded points (aligned with PS), calibrated loss / N, and a key of
    the score vector (hv . s) that identifies equivalent points."""
    P = PNG.shape[0]
    R = np.zeros((P, PS.shape[1]))
    L = np.empty(P)
    K = np.empty(P)
    n = QT.shape[1]
    for t in range(P):
        p = plen[t]
        a0 = GOFF[t]
        ng = PNG[t]
        S = PS[t, :p].copy()
        XS = np.empty((ng, p))
        for g in range(ng):
            for q in range(p):
                XS[g, q] = xval(QT, var_of, tptr, tabv, GREP[a0 + g], S[q])
        r, l = calib_round_kernel(XS, GY[a0:a0 + ng], GC[a0:a0 + ng], PW[t, 1:p + 1].copy(), n_mult, bound)
        R[t, :p] = r
        L[t] = l / N
        sc = score_kernel(QT, var_of, tptr, tabv, S, r)
        kk = 0.0
        for i in range(n):
            kk += hv[i] * sc[i]
        K[t] = kk
    return R, L, K


@_kernel()
def refit_round(w, QT, var_of, lev_of, kind, ncode, tptr, tabv, y, c, bound, tol, n_mult):
    """Continuous logistic fit (box [-bound, bound]) on the support of w, rounded by the calibrated loss over a
    grid of scales (as the beam's final nodes are): a move that changes every point at once."""
    d = w.shape[0]
    n = QT.shape[1]
    cnt = 0
    for j in range(d):
        if w[j] != 0.0:
            cnt += 1
    S = np.empty(cnt, np.int64)
    u = 0
    for j in range(d):
        if w[j] != 0.0:
            S[u] = j
            u += 1
    ginv = np.zeros(n, np.int64)
    code = np.empty(n, np.int64)
    for i in range(n):
        code[i] = 1 if y[i] > 0 else 0
    ginv, ng, gy, gc, rep = regroup(ginv, 1, code, 2, y, c)
    for q in range(cnt):
        j = S[q]
        v = var_of[j]
        if kind[v] == 0:
            l = lev_of[j]
            for i in range(n):
                code[i] = 1 if QT[v, i] >= l else 0
            nc = 2
        else:
            for i in range(n):
                code[i] = QT[v, i]
            nc = ncode[v]
        ginv, ng, gy, gc, rep = regroup(ginv, ng, code, nc, y, c)
    Z = group_design(QT, var_of, tptr, tabv, gy, rep, S)
    npos = 0.0
    tot = 0.0
    for g in range(ng):
        tot += gc[g]
        if gy[g] > 0:
            npos += gc[g]
    beta = np.zeros(cnt + 1)
    beta[0] = np.log(npos / (tot - npos))
    lo = np.full(cnt + 1, -bound)
    hi = np.full(cnt + 1, bound)
    lo[0] = -1e300
    hi[0] = 1e300
    newton_fit(Z, gc, beta, lo, hi, 50, tol)
    XS = np.empty((ng, cnt))
    for g in range(ng):
        for q in range(cnt):
            XS[g, q] = Z[g, q + 1] * gy[g]
    r, l = calib_round_kernel(XS, gy, gc, beta[1:].copy(), n_mult, bound)
    out = np.zeros(d)
    for q in range(cnt):
        out[S[q]] = r[q]
    return out, l


@_kernel()
def bin_scores(s, y, c):
    """Group rows by distinct score: inv (row -> bin), bin scores sv, weights of y=+1 (Wp) and y=-1 (Wn)."""
    n = s.shape[0]
    lo = s[0]
    hi = s[0]
    integral = True
    for i in range(n):
        v = s[i]
        if v < lo:
            lo = v
        if v > hi:
            hi = v
        if integral and v != np.floor(v):
            integral = False
    if integral and hi - lo <= 4 * n + 1024:
        # integer scores in a small range: counting instead of sorting
        R = int(hi - lo) + 1
        cid = np.full(R, -1, np.int64)
        for i in range(n):
            cid[int(s[i] - lo)] = 0
        nb_ = 0
        for r in range(R):
            if cid[r] >= 0:
                cid[r] = nb_
                nb_ += 1
        inv = np.empty(n, np.int64)
        sv = np.empty(nb_)
        Wp = np.zeros(nb_)
        Wn = np.zeros(nb_)
        for r in range(R):
            if cid[r] >= 0:
                sv[cid[r]] = lo + r
        for i in range(n):
            q = cid[int(s[i] - lo)]
            inv[i] = q
            if y[i] > 0:
                Wp[q] += c[i]
            else:
                Wn[q] += c[i]
        return inv, sv, Wp, Wn
    order = np.argsort(s)
    inv = np.empty(n, np.int64)
    sv = np.empty(n)
    Wp = np.zeros(n)
    Wn = np.zeros(n)
    nb_ = -1
    last = 0.0
    for t in range(n):
        i = order[t]
        if t == 0 or s[i] != last:
            nb_ += 1
            sv[nb_] = s[i]
            last = s[i]
        inv[i] = nb_
        if y[i] > 0:
            Wp[nb_] += c[i]
        else:
            Wn[nb_] += c[i]
    nb_ += 1
    return inv, sv[:nb_].copy(), Wp[:nb_].copy(), Wn[:nb_].copy()


@_kernel()
def calibrate_bins(sv, Wp, Wn, a, b):
    """calibrate() on the bins (score sv, weight Wp of y = +1, Wn of y = -1) without building the 2m-row arrays:
    the same terms in the same order (all y = +1 terms, then all y = -1 terms), so the same result."""
    m = sv.shape[0]
    return calibrate_bins_buf(sv, Wp, Wn, a, b, np.empty(2 * m), np.empty(2 * m))


@_kernel()
def calibrate_bins_buf(sv, Wp, Wn, a, b, E, En):
    """calibrate_bins with caller work buffers E, En (length >= 2m)."""
    m = sv.shape[0]
    cur = 0.0
    for i in range(m):
        lv, E[i] = _lrow_e(a * sv[i] + b)
        cur += Wp[i] * lv
    for i in range(m):
        lv, E[m + i] = _lrow_e(-(a * sv[i] + b))
        cur += Wn[i] * lv
    for _ in range(100):
        ga = gb = haa = hab = hbb = 0.0
        for i in range(m):
            z = a * sv[i] + b
            p = _sig_e(z, E[i])
            w = Wp[i] * p * (1.0 - p)
            gi = -Wp[i] * p
            ga += gi * sv[i]
            gb += gi
            haa += w * sv[i] * sv[i]
            hab += w * sv[i]
            hbb += w
        for i in range(m):
            z = -(a * sv[i] + b)
            p = _sig_e(z, E[m + i])
            w = Wn[i] * p * (1.0 - p)
            gi = Wn[i] * p
            ga += gi * sv[i]
            gb += gi
            haa += w * sv[i] * sv[i]
            hab += w * sv[i]
            hbb += w
        haa += 1e-12
        hbb += 1e-12
        det = haa * hbb - hab * hab
        if det <= 1e-18 * (haa * hbb + 1e-300):
            da = 0.0
            db = gb / hbb
        else:
            da = (hbb * ga - hab * gb) / det
            db = (haa * gb - hab * ga) / det
        dec = ga * da + gb * db
        t = 1.0
        new = cur
        while t > 1e-10:
            na, nb_ = a - t * da, b - t * db
            new = 0.0
            for i in range(m):
                lv, En[i] = _lrow_e(na * sv[i] + nb_)
                new += Wp[i] * lv
            for i in range(m):
                lv, En[m + i] = _lrow_e(-(na * sv[i] + nb_))
                new += Wn[i] * lv
            if new <= cur - 1e-4 * t * dec:
                break
            t *= 0.5
        if t <= 1e-10:
            break
        a, b = na, nb_
        E, En = En, E
        improv = cur - new
        cur = new
        if improv < 1e-12 * (1.0 + cur):
            break
    return cur, a, b



@_kernel()
def bin_stats(sv, Wp, Wn, a, b, lp, ln_, pp, pn, wq):
    T = np.zeros(6)
    for q in range(sv.shape[0]):
        z = a * sv[q] + b
        e = np.exp(-abs(z))
        l1 = np.log1p(e)
        if z > 0:
            lp[q] = l1
            ln_[q] = z + l1
            pp[q] = e / (1.0 + e)
            pn[q] = 1.0 / (1.0 + e)
        else:
            lp[q] = -z + l1
            ln_[q] = l1
            pp[q] = 1.0 / (1.0 + e)
            pn[q] = e / (1.0 + e)
        wq[q] = pp[q] * pn[q]
        g = -Wp[q] * pp[q] + Wn[q] * pn[q]
        wt = (Wp[q] + Wn[q]) * wq[q]
        T[0] += Wp[q] * lp[q] + Wn[q] * ln_[q]
        T[1] += g * sv[q]
        T[2] += g
        T[3] += wt * sv[q] * sv[q]
        T[4] += wt * sv[q]
        T[5] += wt
    return T


@_kernel()
def eval_cells(q, deltas, cb, cx, cwp, cwn, nt, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b, tlo, TB):
    """out[:, q, r]: estimated calibrated loss and loss at fixed (a, b) after s += deltas[q, r] * x, from the
    touched cells (bin cb, value cx, weights cwp / cwn of y = +1 / -1). The estimate is the exact loss at the
    current (a, b) minus a trust-region Newton step in (a, b)."""
    nd = deltas.shape[1]
    wk = np.zeros((9, nd))  # one allocation for the nine work rows
    stepa = wk[0]
    stepb = wk[1]
    hyb = wk[2]
    dL = wk[3]
    dGa = wk[4]
    dGb = wk[5]
    dHaa = wk[6]
    dHab = wk[7]
    dHbb = wk[8]
    oL = oGa = oGb = oHaa = oHab = oHbb = 0.0
    for u in range(nt):
        bq = cb[u]
        x = cx[u]
        wpc = cwp[u]
        wnc = cwn[u]
        so = sv[bq]
        lo = wpc * lp[bq] + wnc * ln_[bq]
        go = -wpc * pp[bq] + wnc * pn[bq]
        wo = (wpc + wnc) * wq[bq]
        oL += lo
        oGa += go * so
        oGb += go
        oHaa += wo * so * so
        oHab += wo * so
        oHbb += wo
        for r in range(nd):
            sn_ = so + deltas[q, r] * x
            tf = sn_ - tlo
            if TB.shape[1] > 0 and tf == np.floor(tf) and tf >= 0.0 and tf < TB.shape[1]:
                # integer scores: the loss terms at (a, b) come from a table over the score range
                ti = int(tf)
                lpn = TB[0, ti]
                lnn = TB[1, ti]
                ppn = TB[2, ti]
                pnn = TB[3, ti]
            else:
                z = a * sn_ + b
                e = np.exp(-abs(z))
                l1 = np.log1p(e)
                if z > 0:
                    lpn = l1
                    lnn = z + l1
                    ppn = e / (1.0 + e)
                    pnn = 1.0 / (1.0 + e)
                else:
                    lpn = -z + l1
                    lnn = l1
                    ppn = 1.0 / (1.0 + e)
                    pnn = e / (1.0 + e)
            gn = -wpc * ppn + wnc * pnn
            wn = (wpc + wnc) * ppn * pnn
            dL[r] += wpc * lpn + wnc * lnn - lo
            dGa[r] += gn * sn_ - go * so
            dGb[r] += gn - go
            dHaa[r] += wn * sn_ * sn_ - wo * so * so
            dHab[r] += wn * sn_ - wo * so
            dHbb[r] += wn - wo
    for r in range(nd):
        ga = T[1] + dGa[r]
        gb = T[2] + dGb[r]
        haa = T[3] + dHaa[r] + 1e-12
        hab = T[4] + dHab[r]
        hbb = T[5] + dHbb[r] + 1e-12
        det = haa * hbb - hab * hab
        if det > 1e-14 * haa * hbb:
            da = (hbb * ga - hab * gb) / det
            db = (haa * gb - hab * ga) / det
        else:
            da = 0.0
            db = gb / hbb
        t = 1.0
        if abs(da) * t > tr_a * abs(a) + 1e-300:
            t = tr_a * abs(a) / abs(da)
        if abs(db) * t > tr_b:
            t = tr_b / abs(db)
        stepa[r] = -t * da
        stepb[r] = -t * db
        out[1, q, r] = T[0] + dL[r]
    # the untouched rows at the Newton point, by their quadratic model around (a, b)
    UGa = T[1] - oGa
    UGb = T[2] - oGb
    UHaa = T[3] - oHaa
    UHab = T[4] - oHab
    UHbb = T[5] - oHbb
    for r in range(nd):
        hyb[r] = (T[0] - oL + UGa * stepa[r] + UGb * stepb[r]
                  + 0.5 * (UHaa * stepa[r] * stepa[r] + 2 * UHab * stepa[r] * stepb[r] + UHbb * stepb[r] * stepb[r]))
    # the touched cells exactly at the Newton point
    for u in range(nt):
        so = sv[cb[u]]
        x = cx[u]
        wpc = cwp[u]
        wnc = cwn[u]
        for r in range(nd):
            z = (a + stepa[r]) * (so + deltas[q, r] * x) + b + stepb[r]
            l1 = np.log1p(np.exp(-abs(z)))
            if z > 0:
                hyb[r] += wpc * l1 + wnc * (z + l1)
            else:
                hyb[r] += wpc * (l1 - z) + wnc * l1
    for r in range(nd):
        out[0, q, r] = min(out[1, q, r], hyb[r])


@_kernel()
def eval_vars(cols, qidx, deltas, inv, nbins, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind,
              ncode, tptr, tabv, out, tr_a, tr_b, vrp, vrows):
    """Moves s += deltas[q] * x_cols[q] for columns sorted by variable (qidx: their rows in deltas / out).
    Rows are histogrammed once per variable by (score bin, code); a chain's columns read suffix sums."""
    m = cols.shape[0]
    n = inv.shape[0]
    # table of the per-score loss terms at (a, b) when scores and moves are integers
    tlo = 0.0
    TB = np.zeros((4, 0))
    integral = True
    lo = np.inf
    hi = -np.inf
    for q in range(nbins):
        if sv[q] != np.floor(sv[q]):
            integral = False
        lo = min(lo, sv[q])
        hi = max(hi, sv[q])
    dlo = 0.0
    dhi = 0.0
    for t in range(m):
        for r in range(deltas.shape[1]):
            dv = deltas[qidx[t], r]
            if dv != np.floor(dv):
                integral = False
            dlo = min(dlo, dv)
            dhi = max(dhi, dv)
    if integral and nbins > 0 and hi - lo + dhi - dlo < 4096:
        tlo = lo + dlo
        R = int(hi + dhi - tlo) + 1
        TB = np.empty((4, R))
        for ti in range(R):
            z = a * (tlo + ti) + b
            e = np.exp(-abs(z))
            l1 = np.log1p(e)
            if z > 0:
                TB[0, ti] = l1
                TB[1, ti] = z + l1
                TB[2, ti] = e / (1.0 + e)
                TB[3, ti] = 1.0 / (1.0 + e)
            else:
                TB[0, ti] = -z + l1
                TB[1, ti] = l1
                TB[2, ti] = 1.0 / (1.0 + e)
                TB[3, ti] = e / (1.0 + e)
    u = 0
    while u < m:
        v = var_of[cols[u]]
        u2 = u
        while u2 < m and var_of[cols[u2]] == v:
            u2 += 1
        nc = ncode[v]
        kv = kind[v]
        maxc = nbins if kv == 0 else (nbins * nc if kv == 1 else n)
        cb = np.empty(maxc, np.int64)
        cx = np.empty(maxc)
        cwp = np.empty(maxc)
        cwn = np.empty(maxc)
        if kv < 2:
            HP = np.zeros((nbins, nc))
            HN = np.zeros((nbins, nc))
            for t_ in range(vrp[v], vrp[v + 1]):
                i = vrows[t_]
                if y[i] > 0:
                    HP[inv[i], QT[v, i]] += c[i]
                else:
                    HN[inv[i], QT[v, i]] += c[i]
            if kv == 0:
                for bq in range(nbins):
                    for qq in range(nc - 2, -1, -1):
                        HP[bq, qq] += HP[bq, qq + 1]
                        HN[bq, qq] += HN[bq, qq + 1]
        for t in range(u, u2):
            j = cols[t]
            nt = 0
            if kv == 0:
                l = lev_of[j]
                for bq in range(nbins):
                    if HP[bq, l] > 0.0 or HN[bq, l] > 0.0:
                        cb[nt] = bq
                        cx[nt] = 1.0
                        cwp[nt] = HP[bq, l]
                        cwn[nt] = HN[bq, l]
                        nt += 1
            elif kv == 1:
                for bq in range(nbins):
                    for qq in range(nc):
                        x = tabv[tptr[j] + qq]
                        if x != 0.0 and (HP[bq, qq] > 0.0 or HN[bq, qq] > 0.0):
                            cb[nt] = bq
                            cx[nt] = x
                            cwp[nt] = HP[bq, qq]
                            cwn[nt] = HN[bq, qq]
                            nt += 1
            else:
                for i in range(n):
                    x = tabv[tptr[j] + QT[v, i]]
                    if x != 0.0:
                        cb[nt] = inv[i]
                        cx[nt] = x
                        cwp[nt] = c[i] if y[i] > 0 else 0.0
                        cwn[nt] = c[i] - cwp[nt]
                        nt += 1
            eval_cells(qidx[t], deltas, cb, cx, cwp, cwn, nt, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b, tlo, TB)
        u = u2
    return out


@_kernel()
def removal_state(s, dj, y, c, a, b):
    """Score bins and statistics at fixed (a, b) after s -= dj, with the per-row screening derivatives."""
    n = s.shape[0]
    s2 = s - dj
    inv, sv, Wp, Wn = bin_scores(s2, y, c)
    m = sv.shape[0]
    lp = np.empty(m)
    ln_ = np.empty(m)
    pp = np.empty(m)
    pn = np.empty(m)
    wq = np.empty(m)
    T = bin_stats(sv, Wp, Wn, a, b, lp, ln_, pp, pn, wq)
    r = np.empty(n)
    h = np.empty(n)
    for i in range(n):
        q = inv[i]
        r[i] = (-pp[q] if y[i] > 0 else pn[q]) * c[i]
        h[i] = wq[q] * c[i]
    return s2, inv, sv, Wp, Wn, lp, ln_, pp, pn, wq, T, r, h


@_kernel(inline="always")
def _qbest(G, H, a, vb):
    """min(0, min over integer v in [-vb, vb], v != 0 of u G + u^2 H / 2 with u = a v): the quadratic is convex in v
    (H >= 0), so only the integers next to its minimiser (and -1, 1 around 0, or the ends) can attain it."""
    best = 0.0
    den = a * H
    if den > 0.0 or den < 0.0:
        vs = -G / den
        lo = np.floor(vs)
        if lo < -vb:
            lo = -vb
        if lo > vb:
            lo = vb
        hi = lo + 1.0
        if hi > vb:
            hi = vb
        for v in (lo, hi, -1.0, 1.0):
            if v == 0.0:
                continue
            u = a * v
            f = u * G + 0.5 * u * u * H
            if f < best:
                best = f
    else:
        for v in (-vb, vb):
            u = a * v
            f = u * G + 0.5 * u * u * H
            if f < best:
                best = f
    return best


@_kernel()
def screen_cols(G12, p, cols, a, dl, m):
    """The m columns (ascending) of `cols` with the best second-order score min_delta u G + u^2 H / 2 (u = a d),
    G and H in columns p and p + 1 of G12."""
    nc = cols.shape[0]
    sc = np.empty(nc)
    vb = 0.0
    for r in range(dl.shape[0]):
        vb = max(vb, abs(dl[r]))
    for t in range(nc):
        sc[t] = _qbest(G12[cols[t], p], G12[cols[t], p + 1], a, vb)
    # partial selection of the m smallest (ties: lower index first)
    m = min(m, nc)
    sel = np.empty(m, np.int64)
    cnt = 0
    for t in range(nc):
        v = sc[t]
        if cnt < m:
            u = cnt
            cnt += 1
        elif v < sc[sel[m - 1]]:
            u = m - 1
        else:
            continue
        while u > 0 and sc[sel[u - 1]] > v:
            sel[u] = sel[u - 1]
            u -= 1
        sel[u] = t
    return cols[np.sort(sel)]


@_kernel()
def swap_phase(s, w, S, free, y, c, a, b, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp,
               nzv, n_screen, ne, tr_a, tr_b, cinv, csv):
    """Swap candidates (remove S[q], add a column with a value): per removal, the n_screen best swap-in columns by
    a second-order screen at the removal state, evaluated by the calibrated-loss estimate; returns the ne best per
    removal as arrays (estimate, removed column, added column, value, loss at fixed (a, b))."""
    n = s.shape[0]
    k = S.shape[0]
    d = w.shape[0]
    nv = nzv.shape[0]
    R = np.empty((n, 2 * k))
    INV = np.empty((k, n), np.int64)
    BIN = List()  # per removal: (sv, lp, ln, pp, pn, wq) stacked, and T
    TT = np.empty((k, 6))
    AB = np.empty((k, 2))
    nb0 = csv.shape[0]
    key = np.empty(n, np.int64)
    for q in range(k):
        j = S[q]
        v = var_of[j]
        base = tptr[j]
        if kind[v] == 0:
            # binary column: the removal state's bins are the current bins split by x_j (rows keyed by (bin, x))
            l = lev_of[j]
            pres = np.zeros(2 * nb0, np.bool_)
            Wp2k = np.zeros(2 * nb0)
            Wn2k = np.zeros(2 * nb0)
            for i in range(n):
                kk = 2 * cinv[i] + (1 if QT[v, i] >= l else 0)
                key[i] = kk
                pres[kk] = True
                if y[i] > 0:
                    Wp2k[kk] += c[i]
                else:
                    Wn2k[kk] += c[i]
            npres = 0
            for kk in range(2 * nb0):
                npres += pres[kk]
            ks = np.empty(npres, np.int64)
            ksc = np.empty(npres)
            u = 0
            for kk in range(2 * nb0):
                if pres[kk]:
                    ks[u] = kk
                    ksc[u] = csv[kk >> 1] - w[j] * (kk & 1)
                    u += 1
            o = np.argsort(ksc, kind="mergesort")
            newid = np.empty(2 * nb0, np.int64)
            svt = np.empty(npres)
            m2 = -1
            for t in range(npres):
                if t == 0 or ksc[o[t]] != svt[m2]:
                    m2 += 1
                    svt[m2] = ksc[o[t]]
                newid[ks[o[t]]] = m2
            m2 += 1
            sv = svt[:m2].copy()
            Wp = np.zeros(m2)
            Wn = np.zeros(m2)
            for kk in range(2 * nb0):
                if pres[kk]:
                    Wp[newid[kk]] += Wp2k[kk]
                    Wn[newid[kk]] += Wn2k[kk]
            lp = np.empty(m2)
            ln_ = np.empty(m2)
            pp = np.empty(m2)
            pn = np.empty(m2)
            wq = np.empty(m2)
            AB[q, 0] = a
            AB[q, 1] = b
            T = bin_stats(sv, Wp, Wn, a, b, lp, ln_, pp, pn, wq)
            inv = np.empty(n, np.int64)
            for i in range(n):
                g = newid[key[i]]
                inv[i] = g
                R[i, 2 * q] = (-pp[g] if y[i] > 0 else pn[g]) * c[i]
                R[i, 2 * q + 1] = wq[g] * c[i]
        else:
            dj = np.empty(n)
            for i in range(n):
                dj[i] = w[j] * tabv[base + QT[v, i]]
            s2, inv, sv, Wp, Wn, lp, ln_, pp, pn, wq, T, r_, h_ = removal_state(s, dj, y, c, a, b)
            AB[q, 0] = a
            AB[q, 1] = b
            R[:, 2 * q] = r_
            R[:, 2 * q + 1] = h_
        INV[q] = inv
        st6 = np.empty((6, sv.shape[0]))
        st6[0] = sv
        st6[1] = lp
        st6[2] = ln_
        st6[3] = pp
        st6[4] = pn
        st6[5] = wq
        BIN.append(st6)
        TT[q] = T
    pow2 = np.zeros(2 * k, np.bool_)
    for q in range(k):
        pow2[2 * q + 1] = True
    G12 = colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pow2, vrp, vrows)
    o_est = np.full(k * ne, np.inf)
    o_fx = np.full(k * ne, np.inf)
    o_rj = np.full(k * ne, -1, np.int64)
    o_aj = np.full(k * ne, -1, np.int64)
    o_v = np.zeros(k * ne)
    for q in range(k):
        j = S[q]
        cnt = 0
        for t in range(d):
            if free[t]:
                cnt += 1
        if cnt == 0:
            continue
        nonSq = np.empty(cnt, np.int64)
        cnt = 0
        for t in range(d):
            if free[t]:
                nonSq[cnt] = t
                cnt += 1
        aq = AB[q, 0]
        bq = AB[q, 1]
        if abs(a) * (csv[csv.shape[0] - 1] - csv[0]) >= SWNG:
            cols = screen_newton(G12, 2 * q, nonSq, n_screen)  # integer steps overshoot at a steep map
        else:
            cols = screen_cols(G12, 2 * q, nonSq, aq, nzv, n_screen)
        m = cols.shape[0]
        st6 = BIN[q]
        inv = INV[q]
        sv = st6[0]
        lp = st6[1]
        ln_ = st6[2]
        pp = st6[3]
        pn = st6[4]
        wq = st6[5]
        T = TT[q]
        vo = np.empty(m, np.int64)
        for t in range(m):
            vo[t] = var_of[cols[t]]
        o = np.argsort(vo, kind="mergesort")
        deltas = np.empty((m, nv))
        for t in range(m):
            deltas[t, :] = nzv
        out = np.empty((2, m, nv))
        eval_vars(cols[o], o, deltas, inv, sv.shape[0], sv, lp, ln_, pp, pn, wq, T, aq, bq, y, c, QT, var_of,
                  lev_of, kind, ncode, tptr, tabv, out, tr_a, tr_b, vrp, vrows)
        flat = out[0].ravel()
        fo = np.argsort(flat, kind="mergesort")
        u = 0
        for t in range(fo.shape[0]):
            if u >= ne:
                break
            f = fo[t]
            if not np.isfinite(flat[f]):
                continue
            r = f % nv
            qq = f // nv
            o_est[q * ne + u] = flat[f]
            o_fx[q * ne + u] = out[1].ravel()[f]
            o_rj[q * ne + u] = j
            o_aj[q * ne + u] = cols[qq]
            o_v[q * ne + u] = nzv[r]
            u += 1
    return o_est, o_rj, o_aj, o_v, o_fx


@_kernel()
def support_hist(S, inv, nb_, y, c, QT, var_of, kind, ncode, lev_of, radius):
    """For every support column of a chain: weights of y = +1 / -1 per (score bin, code of its variable), suffix-
    summed over codes (entry [q, b, l]: rows of bin b with code >= l), from one pass over the rows."""
    k = S.shape[0]
    n = inv.shape[0]
    vs = np.empty(k, np.int64)
    mc = 1
    for q in range(k):
        vs[q] = var_of[S[q]]
        if kind[vs[q]] == 0:
            mc = max(mc, ncode[vs[q]] + 1)
    HP = np.zeros((k, nb_, mc))
    HN = np.zeros((k, nb_, mc))
    for q in range(k):
        if kind[vs[q]] != 0:
            vs[q] = -1
    for i in range(n):
        bq = inv[i]
        ci = c[i]
        if y[i] > 0:
            for q in range(k):
                if vs[q] >= 0:
                    HP[q, bq, QT[vs[q], i]] += ci
        else:
            for q in range(k):
                if vs[q] >= 0:
                    HN[q, bq, QT[vs[q], i]] += ci
    for q in range(k):
        if vs[q] < 0:
            continue
        nc = ncode[vs[q]]
        # suffix sums down to the lowest level a value change or a slide reads
        lo_q = max(0, lev_of[S[q]] - radius)
        for bq in range(nb_):
            for qq in range(nc - 1, lo_q - 1, -1):
                HP[q, bq, qq] += HP[q, bq, qq + 1]
                HN[q, bq, qq] += HN[q, bq, qq + 1]
    return HP, HN


@_kernel()
def make_tb(sv, nbins, dlo, dhi, a, b):
    """Loss terms at (a, b) tabulated over the integer scores sv + [dlo, dhi] (empty when scores are not integers)."""
    tlo = 0.0
    TB = np.zeros((4, 0))
    integral = True
    lo = np.inf
    hi = -np.inf
    for q in range(nbins):
        if sv[q] != np.floor(sv[q]):
            integral = False
        lo = min(lo, sv[q])
        hi = max(hi, sv[q])
    if dlo != np.floor(dlo) or dhi != np.floor(dhi):
        integral = False
    if integral and nbins > 0 and hi - lo + dhi - dlo < 4096:
        tlo = lo + dlo
        R = int(hi + dhi - tlo) + 1
        TB = np.empty((4, R))
        for ti in range(R):
            z = a * (tlo + ti) + b
            e = np.exp(-abs(z))
            l1 = np.log1p(e)
            if z > 0:
                TB[0, ti] = l1
                TB[1, ti] = z + l1
                TB[2, ti] = e / (1.0 + e)
                TB[3, ti] = 1.0 / (1.0 + e)
            else:
                TB[0, ti] = -z + l1
                TB[1, ti] = l1
                TB[2, ti] = 1.0 / (1.0 + e)
                TB[3, ti] = e / (1.0 + e)
    return tlo, TB


@_kernel()
def eval_support(S, deltas, HP, HN, lev_of, nb_, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b):
    """Value changes of the support columns (all chain columns) from the support histograms."""
    m = S.shape[0]
    dlo = 0.0
    dhi = 0.0
    integral = True
    for q in range(m):
        for r in range(deltas.shape[1]):
            dv = deltas[q, r]
            if dv != np.floor(dv):
                integral = False
            dlo = min(dlo, dv)
            dhi = max(dhi, dv)
    if integral:
        tlo, TB = make_tb(sv, nb_, dlo, dhi, a, b)
    else:
        tlo, TB = 0.0, np.zeros((4, 0))
    cb = np.empty(nb_, np.int64)
    cx = np.ones(nb_)
    cwp = np.empty(nb_)
    cwn = np.empty(nb_)
    for q in range(m):
        l = lev_of[S[q]]
        nt = 0
        for bq in range(nb_):
            if HP[q, bq, l] > 0.0 or HN[q, bq, l] > 0.0:
                cb[nt] = bq
                cwp[nt] = HP[q, bq, l]
                cwn[nt] = HN[q, bq, l]
                nt += 1
        eval_cells(q, deltas, cb, cx, cwp, cwn, nt, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b, tlo, TB)


@_kernel()
def slide_moves(w, S, free, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind, ncode, vrp,
                vrows, vcols, vcp, ne, tr_a, tr_b, radius, SHP, SHN):
    """Threshold slides: a support column of a chain moves to another threshold of its variable (at most `radius`
    levels away) with the same points. The score changes by +-w_j on the band of codes between the two levels,
    so every slide of a column is estimated from one (bin, code) histogram of its variable. Returns the ne best
    per support column (estimate, removed column, added column, value)."""
    m = S.shape[0]
    nb_ = sv.shape[0]
    o_est = np.full(m * ne, np.inf)
    o_rj = np.full(m * ne, -1, np.int64)
    o_aj = np.full(m * ne, -1, np.int64)
    o_v = np.zeros(m * ne)
    cb = np.empty(nb_, np.int64)
    cx = np.empty(nb_)
    cwp = np.empty(nb_)
    cwn = np.empty(nb_)
    for q in range(m):
        j = S[q]
        v = var_of[j]
        nc = ncode[v]
        if kind[v] != 0 or nc <= 2:
            continue
        l = lev_of[j]
        lo_l = max(1, l - radius)
        hi_l = min(nc - 1, l + radius)
        ntar = 0
        tl = np.empty(hi_l - lo_l + 1, np.int64)
        for l2 in range(lo_l, hi_l + 1):
            j2 = vcols[vcp[v] + l2 - 1]
            if l2 != l and free[j2]:
                tl[ntar] = l2
                ntar += 1
        if ntar == 0:
            continue
        HP = SHP[q]
        HN = SHN[q]
        deltas = np.full((ntar, 1), w[j])
        tlo, TB = make_tb(sv, nb_, -abs(w[j]), abs(w[j]), a, b)
        out = np.empty((2, ntar, 1))
        for t in range(ntar):
            l2 = tl[t]
            lo2 = min(l, l2)
            hi2 = max(l, l2)
            sg = 1.0 if l2 < l else -1.0  # rows of the band gain (lower level) or lose the indicator
            nt = 0
            for bq in range(nb_):
                wp_ = HP[bq, lo2] - HP[bq, hi2]
                wn_ = HN[bq, lo2] - HN[bq, hi2]
                if wp_ > 0.0 or wn_ > 0.0:
                    cb[nt] = bq
                    cx[nt] = sg
                    cwp[nt] = wp_
                    cwn[nt] = wn_
                    nt += 1
            eval_cells(t, deltas, cb, cx, cwp, cwn, nt, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b, tlo, TB)
        fo = np.argsort(out[0, :, 0], kind="mergesort")
        for u in range(min(ne, ntar)):
            t = fo[u]
            o_est[q * ne + u] = out[0, t, 0]
            o_rj[q * ne + u] = j
            o_aj[q * ne + u] = vcols[vcp[v] + tl[t] - 1]
            o_v[q * ne + u] = w[j]
    return o_est, o_rj, o_aj, o_v


@_kernel()
def main_phase(w, S, free, k, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind, ncode, tptr,
               tabv, vrp, vrows, vcols, vcp, allv, nzv, ne, tr_a, tr_b, add_screen, SHP, SHN, Wp, Wn):
    """Value changes of the support columns (every other value; 0 removes) and additions (when |S| < k): the ne
    best by the estimate per support column, and the ne best additions overall."""
    m = S.shape[0]
    d = w.shape[0]
    nv = nzv.shape[0]
    nb_ = sv.shape[0]
    o_est = np.full(m * ne + ne, np.inf)
    o_fx = np.full(m * ne + ne, np.inf)
    o_aj = np.full(m * ne + ne, -1, np.int64)
    o_v = np.zeros(m * ne + ne)
    if m > 0:
        newv = np.empty((m, allv.shape[0] - 1))
        deltas = np.empty((m, allv.shape[0] - 1))
        vo = np.empty(m, np.int64)
        for q in range(m):
            u = 0
            wj = w[S[q]]
            for r in range(allv.shape[0]):
                if allv[r] != wj:
                    newv[q, u] = allv[r]
                    deltas[q, u] = allv[r] - wj
                    u += 1
            vo[q] = var_of[S[q]]
        o = np.argsort(vo, kind="mergesort")
        out = np.empty((2, m, newv.shape[1]))
        allch = True
        for q in range(m):
            if kind[var_of[S[q]]] != 0:
                allch = False
        if allch:
            eval_support(S, deltas, SHP, SHN, lev_of, nb_, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b)
        else:
            eval_vars(S[o], o, deltas, inv, nb_, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind,
                      ncode, tptr, tabv, out, tr_a, tr_b, vrp, vrows)
        for q in range(m):
            fo = np.argsort(out[0, q], kind="mergesort")
            u = 0
            for t in range(fo.shape[0]):
                if u >= ne:
                    break
                r = fo[t]
                if not np.isfinite(out[0, q, r]):
                    continue
                o_est[q * ne + u] = out[0, q, r]
                o_fx[q * ne + u] = out[1, q, r]
                o_aj[q * ne + u] = S[q]
                o_v[q * ne + u] = newv[q, r]
                u += 1
    if m < k:
        cnt = 0
        for t in range(d):
            if free[t]:
                cnt += 1
        if cnt > 0:
            cols = np.empty(cnt, np.int64)
            vo = np.empty(cnt, np.int64)
            cnt = 0
            for t in range(d):
                if free[t]:
                    cols[cnt] = t
                    vo[cnt] = var_of[t]
                    cnt += 1
            if add_screen > 0 and cnt > add_screen:
                # second-order screen of the additions at the current (a, b); only the best are estimated
                R2 = np.empty((y.shape[0], 2))
                for i in range(y.shape[0]):
                    q = inv[i]
                    R2[i, 0] = (-pp[q] if y[i] > 0 else pn[q]) * c[i]
                    R2[i, 1] = wq[q] * c[i]
                pw2 = np.zeros(2, np.bool_)
                pw2[1] = True
                G12 = colsum(QT, R2, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pw2, vrp, vrows)
                cols = screen_cols(G12, 0, cols, a, nzv, add_screen)
                cnt = cols.shape[0]
                vo = np.empty(cnt, np.int64)
                for t in range(cnt):
                    vo[t] = var_of[cols[t]]
            o = np.argsort(vo, kind="mergesort")
            deltas = np.empty((cnt, nv))
            for t in range(cnt):
                deltas[t, :] = nzv
            out = np.empty((2, cnt, nv))
            eval_vars(cols[o], o, deltas, inv, nb_, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of,
                      kind, ncode, tptr, tabv, out, tr_a, tr_b, vrp, vrows)
            flat = out[0].ravel()
            fl1 = out[1].ravel()
            # the ne smallest finite estimates
            u = 0
            for _ in range(ne):
                bi = -1
                bv = np.inf
                for f in range(flat.shape[0]):
                    if flat[f] < bv:
                        dup = False
                        for z in range(u):
                            if o_aj[m * ne + z] == cols[f // nv] and o_v[m * ne + z] == nzv[f % nv]:
                                dup = True
                        if not dup:
                            bv = flat[f]
                            bi = f
                if bi < 0:
                    break
                o_est[m * ne + u] = flat[bi]
                o_fx[m * ne + u] = fl1[bi]
                o_aj[m * ne + u] = cols[bi // nv]
                o_v[m * ne + u] = nzv[bi % nv]
                u += 1
    return o_est, o_aj, o_v, o_fx


@_kernel()
def check_kernel(s, w, L, est, rjs, ajs, vs, a, b, margin, y, c, QT, var_of, tptr, tabv, kind, lev_of, cinv, csv):
    """Exact calibrated loss of the candidate moves (in order); returns the index of the best improving one (-1 if
    none) and its score vector, bins and calibration. Moves on binary columns are scored from rows keyed by
    (score bin, x_removed, x_added) in one pass."""
    n = s.shape[0]
    bi = -1
    bL = np.inf
    ba = a
    bb = b
    nb0 = csv.shape[0]
    Wp4 = np.zeros(4 * nb0)
    Wn4 = np.zeros(4 * nb0)
    sc = np.empty(4 * nb0)
    kp = np.empty(4 * nb0, np.int64)
    for t in range(est.shape[0]):
        if est[t] > L * (1.0 + margin):
            continue  # the estimate says the move does not help
        rj = rjs[t]
        aj = ajs[t]
        va = var_of[aj]
        if kind[va] == 0 and (rj < 0 or kind[var_of[rj]] == 0):
            la = lev_of[aj]
            dv = vs[t] - w[aj]
            vr = var_of[rj] if rj >= 0 else 0
            lr = lev_of[rj] if rj >= 0 else 0
            wr = w[rj] if rj >= 0 else 0.0
            Wp4[:] = 0.0
            Wn4[:] = 0.0
            for i in range(n):
                kk = 4 * cinv[i] + (1 if QT[va, i] >= la else 0)
                if rj >= 0 and QT[vr, i] >= lr:
                    kk += 2
                if y[i] > 0:
                    Wp4[kk] += c[i]
                else:
                    Wn4[kk] += c[i]
            m = 0
            for kk in range(4 * nb0):
                if Wp4[kk] > 0.0 or Wn4[kk] > 0.0:
                    kp[m] = kk
                    sc[m] = csv[kk >> 2] - wr * ((kk >> 1) & 1) + dv * (kk & 1)
                    m += 1
            o = np.argsort(sc[:m], kind="mergesort")
            sv = np.empty(m)
            Wp = np.zeros(m)
            Wn = np.zeros(m)
            g = -1
            for u in range(m):
                kk = kp[o[u]]
                if g < 0 or sc[o[u]] != sv[g]:
                    g += 1
                    sv[g] = sc[o[u]]
                Wp[g] += Wp4[kk]
                Wn[g] += Wn4[kk]
            g += 1
            if g < 2:
                continue
            L2, a2, b2 = calibrate_bins(sv[:g], Wp[:g], Wn[:g], a, b)
        else:
            s2 = s.copy()
            if rj >= 0:
                v = var_of[rj]
                base = tptr[rj]
                for i in range(n):
                    s2[i] -= w[rj] * tabv[base + QT[v, i]]
            dv = vs[t] - w[aj]
            base = tptr[aj]
            for i in range(n):
                s2[i] += dv * tabv[base + QT[va, i]]
            lo = s2[0]
            hi = s2[0]
            for i in range(n):
                lo = min(lo, s2[i])
                hi = max(hi, s2[i])
            if hi == lo:
                continue
            inv, sv, Wp, Wn = bin_scores(s2, y, c)
            L2, a2, b2 = calibrate_bins(sv, Wp, Wn, a, b)
        if L2 < L - 1e-9 * L and L2 < bL:
            bi, bL, ba, bb = t, L2, a2, b2
            break  # first improving candidate in estimate order
    if bi < 0:
        return bi, bL, s, np.zeros(0, np.int64), np.zeros(0), np.zeros(0), np.zeros(0), a, b
    # the winner's score vector and bins
    s2 = s.copy()
    rj = rjs[bi]
    if rj >= 0:
        v = var_of[rj]
        base = tptr[rj]
        for i in range(n):
            s2[i] -= w[rj] * tabv[base + QT[v, i]]
    aj = ajs[bi]
    dv = vs[bi] - w[aj]
    v = var_of[aj]
    base = tptr[aj]
    for i in range(n):
        s2[i] += dv * tabv[base + QT[v, i]]
    inv, sv, Wp, Wn = bin_scores(s2, y, c)
    return bi, bL, s2, inv, sv, Wp, Wn, ba, bb


@_kernel()
def exact_values(w, S, s, L, a, b, inv, sv, Wp, Wn, y, c, QT, var_of, kind, ncode, lev_of, tptr, tabv, allv):
    """Exact calibrated loss of every value change of every support column (chain columns from one (bin, code)
    histogram pass; other columns by a row pass each). Returns (best total loss, column, value); column -1 if none
    improves."""
    nb_ = sv.shape[0]
    SHP, SHN = support_hist(S, inv, nb_, y, c, QT, var_of, kind, ncode, lev_of, 0)
    bL = L - 1e-9 * L
    bj = -1
    bv = 0.0
    m2 = 2 * nb_
    cwp = np.empty(m2)
    cwn = np.empty(m2)
    svq = np.empty(m2)
    Wpq = np.empty(m2)
    Wnq = np.empty(m2)
    E = np.empty(2 * m2)
    En = np.empty(2 * m2)
    n = s.shape[0]
    for q in range(S.shape[0]):
        j = S[q]
        v = var_of[j]
        if kind[v] == 0:
            l = lev_of[j]
            for bq in range(nb_):
                cwp[2 * bq] = Wp[bq] - SHP[q, bq, l]
                cwn[2 * bq] = Wn[bq] - SHN[q, bq, l]
                cwp[2 * bq + 1] = SHP[q, bq, l]
                cwn[2 * bq + 1] = SHN[q, bq, l]
            c0p = np.empty(nb_)
            c0n = np.empty(nb_)
            c1p = np.empty(nb_)
            c1n = np.empty(nb_)
            for bq in range(nb_):
                c0p[bq] = cwp[2 * bq]
                c0n[bq] = cwn[2 * bq]
                c1p[bq] = cwp[2 * bq + 1]
                c1n[bq] = cwn[2 * bq + 1]
            # loss at the current (a, b) of every value; the POLVV best calibrated exactly
            fl = np.full(allv.shape[0], np.inf)
            for r in range(allv.shape[0]):
                dv = allv[r] - w[j]
                if dv == 0.0:
                    continue
                t = 0.0
                for bq in range(nb_):
                    if c1p[bq] <= 0.0 and c1n[bq] <= 0.0:
                        continue
                    z = a * (sv[bq] + dv) + b
                    l1 = np.log1p(np.exp(-abs(z)))
                    if z > 0:
                        t += c1p[bq] * l1 + c1n[bq] * (z + l1)
                    else:
                        t += c1p[bq] * (l1 - z) + c1n[bq] * l1
                fl[r] = t
            fo = np.argsort(fl)
            for r0 in range(min(POLVV, allv.shape[0] - 1)):
                r = fo[r0]
                dv = allv[r] - w[j]
                L2 = _merge_calib(sv, c0p, c0n, c1p, c1n, dv, a, b, svq, Wpq, Wnq, E, En, bL)
                if L2 < bL:
                    bL, bj, bv = L2, j, allv[r]
        else:
            base = tptr[j]
            s2 = np.empty(n)
            for r in range(allv.shape[0]):
                dv = allv[r] - w[j]
                if dv == 0.0:
                    continue
                for i in range(n):
                    s2[i] = s[i] + dv * tabv[base + QT[v, i]]
                inv2, sv2, Wp2, Wn2 = bin_scores(s2, y, c)
                if sv2.shape[0] < 2:
                    continue
                L2, a2, b2 = calibrate_bins(sv2, Wp2, Wn2, a, b)
                if L2 < bL:
                    bL, bj, bv = L2, j, allv[r]
    return bL, bj, bv


@_kernel()
def calib_target(sv, Wp, Wn, a, b, target, E, En):
    """calibrate_bins_buf that stops early: returns the loss once it is below target (an upper bound of the
    calibrated loss, so the move surely beats target), or np.inf once even twice the Newton decrement cannot reach
    target (rejected), else the converged loss."""
    m = sv.shape[0]
    cur = 0.0
    for i in range(m):
        lv, E[i] = _lrow_e(a * sv[i] + b)
        cur += Wp[i] * lv
    for i in range(m):
        lv, E[m + i] = _lrow_e(-(a * sv[i] + b))
        cur += Wn[i] * lv
    if cur < target:
        return cur
    for _ in range(CALCAP):
        ga = gb = haa = hab = hbb = 0.0
        for i in range(m):
            z = a * sv[i] + b
            p = _sig_e(z, E[i])
            w = Wp[i] * p * (1.0 - p)
            gi = -Wp[i] * p
            ga += gi * sv[i]
            gb += gi
            haa += w * sv[i] * sv[i]
            hab += w * sv[i]
            hbb += w
        for i in range(m):
            z = -(a * sv[i] + b)
            p = _sig_e(z, E[m + i])
            w = Wn[i] * p * (1.0 - p)
            gi = Wn[i] * p
            ga += gi * sv[i]
            gb += gi
            haa += w * sv[i] * sv[i]
            hab += w * sv[i]
            hbb += w
        haa += 1e-12
        hbb += 1e-12
        det = haa * hbb - hab * hab
        if det <= 1e-18 * (haa * hbb + 1e-300):
            da = 0.0
            db = gb / hbb
        else:
            da = (hbb * ga - hab * gb) / det
            db = (haa * gb - hab * ga) / det
        dec = ga * da + gb * db
        if cur - REJ * dec > target:
            return np.inf
        t = 1.0
        new = cur
        while t > 1e-10:
            na, nb_ = a - t * da, b - t * db
            new = 0.0
            for i in range(m):
                lv, En[i] = _lrow_e(na * sv[i] + nb_)
                new += Wp[i] * lv
            for i in range(m):
                lv, En[m + i] = _lrow_e(-(na * sv[i] + nb_))
                new += Wn[i] * lv
            if new <= cur - 1e-4 * t * dec:
                break
            t *= 0.5
        if t <= 1e-10:
            break
        a, b = na, nb_
        E, En = En, E
        improv = cur - new
        cur = new
        if cur < target:
            return cur
        if improv < 1e-12 * (1.0 + cur):
            break
    return cur


@_kernel()
def _pav_loss(Wp, Wn, g, sgn, bp, bn):
    """Log loss of the best monotone (sgn > 0: non-decreasing, else non-increasing) probabilities over cells in
    score order (pool adjacent violators); a lower bound of the loss of every logistic map with that slope sign."""
    nbk = 0
    for u0 in range(g):
        u = u0 if sgn > 0 else g - 1 - u0
        bp[nbk] = Wp[u]
        bn[nbk] = Wn[u]
        nbk += 1
        while nbk > 1 and bp[nbk - 2] * (bp[nbk - 1] + bn[nbk - 1]) >= bp[nbk - 1] * (bp[nbk - 2] + bn[nbk - 2]):
            bp[nbk - 2] += bp[nbk - 1]
            bn[nbk - 2] += bn[nbk - 1]
            nbk -= 1
    L = 0.0
    for t in range(nbk):
        tot = bp[t] + bn[t]
        if bp[t] > 0.0:
            L -= bp[t] * np.log(bp[t] / tot)
        if bn[t] > 0.0:
            L -= bn[t] * np.log(bn[t] / tot)
    return L


@_kernel()
def _merge_calib(sv, cwp0, cwn0, cwp1, cwn1, dv, a, b, svq, Wpq, Wnq, E, En, target):
    """Exact calibrated loss of the cells (score sv[q], weights cw*0) and (score sv[q] + dv, weights cw*1), sv
    ascending: merge of the two sorted lists, equal scores combined."""
    m = sv.shape[0]
    i0 = 0
    i1 = 0
    g = -1
    while i0 < m or i1 < m:
        if i1 >= m or (i0 < m and sv[i0] <= sv[i1] + dv):
            sc = sv[i0]
            wp = cwp0[i0]
            wn = cwn0[i0]
            i0 += 1
        else:
            sc = sv[i1] + dv
            wp = cwp1[i1]
            wn = cwn1[i1]
            i1 += 1
        if wp <= 0.0 and wn <= 0.0:
            continue
        if g < 0 or sc != svq[g]:
            g += 1
            svq[g] = sc
            Wpq[g] = 0.0
            Wnq[g] = 0.0
        Wpq[g] += wp
        Wnq[g] += wn
    g += 1
    if g < 2:
        return np.inf
    # exact rejection: the best monotone map's loss bounds every calibrated loss from below
    lb = min(_pav_loss(Wpq, Wnq, g, 1, E, En), _pav_loss(Wpq, Wnq, g, -1, E, En))
    if lb >= target:
        return np.inf
    return calib_target(svq[:g], Wpq[:g], Wnq[:g], a, b, target, E, En)


@_kernel()
def screen_newton(G12, p, cols, m):
    """The m columns of `cols` with the largest continuous Newton gain G^2 / H (scale free: the refitted map can
    absorb any step size), ascending."""
    nc = cols.shape[0]
    sc = np.empty(nc)
    for t in range(nc):
        g = G12[cols[t], p]
        h = G12[cols[t], p + 1]
        sc[t] = -g * g / h if h > 1e-300 else 0.0
    m = min(m, nc)
    sel = np.empty(m, np.int64)
    cnt = 0
    for t in range(nc):
        v = sc[t]
        if cnt < m:
            u = cnt
            cnt += 1
        elif v < sc[sel[m - 1]]:
            u = m - 1
        else:
            continue
        while u > 0 and sc[sel[u - 1]] > v:
            sel[u] = sel[u - 1]
            u -= 1
        sel[u] = t
    return cols[np.sort(sel)]


@_kernel()
def exact_moves(s, w, S, free, k, y, c, a, b, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp,
                nzv, n_screen, L0):
    """Exact best swap / addition: per removal (and no removal when |S| < k) the n_screen best columns by the
    continuous Newton gain at the removal state with a refitted (a, b) (one colsum for all removals); the POLV best
    values of every screened chain column (by the loss at the refitted (a, b)) scored
    by their exact calibrated loss on (bin, x) cells. Returns (total loss, removed column or -1, added column,
    value); added column -1 when nothing beats L0."""
    n = s.shape[0]
    m = S.shape[0]
    d = w.shape[0]
    nv = nzv.shape[0]
    vb = 0.0
    for r in range(nv):
        vb = max(vb, abs(nzv[r]))
    bL = L0
    brj = -1
    baj = -1
    bv = 0.0
    nfree = 0
    for t in range(d):
        if free[t] and kind[var_of[t]] == 0:
            nfree += 1
    if nfree == 0:
        return bL, brj, baj, bv
    fcols = np.empty(nfree, np.int64)
    u = 0
    for t in range(d):
        if free[t] and kind[var_of[t]] == 0:
            fcols[u] = t
            u += 1
    q0 = -1 if m < k else 0
    nq = m - q0
    R = np.empty((n, 2 * nq))
    INV = np.empty((nq, n), np.int64)
    BINS = List()
    AB = np.empty((nq, 2))
    s2 = np.empty(n)
    for qi in range(nq):
        q = q0 + qi
        if q >= 0:
            j = S[q]
            base = tptr[j]
            vj = var_of[j]
            for i in range(n):
                s2[i] = s[i] - w[j] * tabv[base + QT[vj, i]]
        else:
            for i in range(n):
                s2[i] = s[i]
        inv, sv, Wp, Wn = bin_scores(s2, y, c)
        nb_ = sv.shape[0]
        if nb_ >= 2:
            _, aq, bq = calibrate_bins(sv, Wp, Wn, a, b)
        else:
            aq, bq = a, b
        AB[qi, 0] = aq
        AB[qi, 1] = bq
        INV[qi] = inv
        st3 = np.empty((3, nb_))
        st3[0] = sv
        st3[1] = Wp
        st3[2] = Wn
        BINS.append(st3)
        lp = np.empty(nb_)
        ln_ = np.empty(nb_)
        pp = np.empty(nb_)
        pn = np.empty(nb_)
        wq = np.empty(nb_)
        bin_stats(sv, Wp, Wn, aq, bq, lp, ln_, pp, pn, wq)
        for i in range(n):
            g = inv[i]
            R[i, 2 * qi] = (-pp[g] if y[i] > 0 else pn[g]) * c[i]
            R[i, 2 * qi + 1] = wq[g] * c[i]
    pw = np.zeros(2 * nq, np.bool_)
    for qi in range(nq):
        pw[2 * qi + 1] = True
    G = colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pw, vrp, vrows)
    fl = np.zeros(nv)
    for qi in range(nq):
        q = q0 + qi
        j = S[q] if q >= 0 else -1
        aq = AB[qi, 0]
        bq = AB[qi, 1]
        st3 = BINS[qi]
        sv = st3[0]
        Wp = st3[1]
        Wn = st3[2]
        nb_ = sv.shape[0]
        inv = INV[qi]
        cols = screen_newton(G, 2 * qi, fcols, n_screen)
        cwp0 = np.empty(nb_)
        cwn0 = np.empty(nb_)
        cwp1 = np.empty(nb_)
        cwn1 = np.empty(nb_)
        svq = np.empty(2 * nb_)
        Wpq = np.empty(2 * nb_)
        Wnq = np.empty(2 * nb_)
        E = np.empty(4 * nb_)
        En = np.empty(4 * nb_)
        tlo, TB = make_tb(sv, nb_, -vb, vb, aq, bq)
        # columns grouped by variable: one (bin, code) histogram per variable, suffix-summed over codes
        vo = np.empty(cols.shape[0], np.int64)
        for t in range(cols.shape[0]):
            vo[t] = var_of[cols[t]] * (d + 1) + lev_of[cols[t]]
        cols = cols[np.argsort(vo)]
        cur_v = -1
        HPv = np.zeros((1, 1))
        HNv = np.zeros((1, 1))
        for t in range(cols.shape[0]):
            l = cols[t]
            if l == j:
                continue
            v = var_of[l]
            lv = lev_of[l]
            if v != cur_v:
                cur_v = v
                nc = ncode[v]
                HPv = np.zeros((nc + 1, nb_))
                HNv = np.zeros((nc + 1, nb_))
                for tt in range(vrp[v], vrp[v + 1]):
                    i = vrows[tt]
                    if y[i] > 0:
                        HPv[QT[v, i], inv[i]] += c[i]
                    else:
                        HNv[QT[v, i], inv[i]] += c[i]
                for cc in range(nc - 1, 0, -1):
                    for g in range(nb_):
                        HPv[cc, g] += HPv[cc + 1, g]
                        HNv[cc, g] += HNv[cc + 1, g]
            for g in range(nb_):
                cwp1[g] = HPv[lv, g]
                cwn1[g] = HNv[lv, g]
            for g in range(nb_):
                cwp0[g] = Wp[g] - cwp1[g]
                cwn0[g] = Wn[g] - cwn1[g]
            fl[:] = 0.0
            for g in range(nb_):
                if cwp1[g] <= 0.0 and cwn1[g] <= 0.0:
                    continue
                for r in range(nv):
                    tf = sv[g] + nzv[r] - tlo
                    if TB.shape[1] > 0 and tf >= 0.0 and tf < TB.shape[1] and tf == np.floor(tf):
                        ti = int(tf)
                        fl[r] += cwp1[g] * TB[0, ti] + cwn1[g] * TB[1, ti]
                        continue
                    z = aq * (sv[g] + nzv[r]) + bq
                    l1 = np.log1p(np.exp(-abs(z)))
                    if z > 0:
                        fl[r] += cwp1[g] * l1 + cwn1[g] * (z + l1)
                    else:
                        fl[r] += cwp1[g] * (l1 - z) + cwn1[g] * l1
            fo = np.argsort(fl)
            for r0 in range(min(POLV, nv)):
                r = fo[r0]
                L2 = _merge_calib(sv, cwp0, cwn0, cwp1, cwn1, nzv[r], aq, bq, svq, Wpq, Wnq, E, En, bL)
                if L2 < bL:
                    bL, brj, baj, bv = L2, j, l, nzv[r]
    return bL, brj, baj, bv


@_kernel()
def top_order(est, m):
    """Indices of the m smallest finite entries of est, ties in index order."""
    o = np.argsort(est, kind="mergesort")
    cnt = 0
    for t in range(o.shape[0]):
        if np.isfinite(est[o[t]]):
            cnt += 1
    cnt = min(cnt, m)
    return o[:cnt]


@_kernel()
def ils_kernel(w, s, L, a, b, inv, sv, Wp, Wn, visited, hw, k, max_iter, swaps, gate, N, ne, n_screen, add_screen,
               margin, y, c, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp, valid,
               allv, nzv, tr_a, tr_b):
    """Best-improvement local search on the integer lattice (value changes and additions; swaps when those fail),
    candidates ranked by the calibrated-loss estimate and the best 2 ne checked exactly. visited: keys of points
    already expanded in this fit (the search from them is deterministic). Returns (loss, w)."""
    d = w.shape[0]
    m_chk = 2 * ne
    # the candidates of a failing swap phase at the final point (reused by the refit swaps)
    has_c = False
    c_oe = np.zeros(0)
    c_orj = np.zeros(0, np.int64)
    c_oaj = np.zeros(0, np.int64)
    c_ov = np.zeros(0)
    for it in range(max_iter):
        key = 0.0
        for j in range(d):
            if w[j] != 0.0:
                key += w[j] * hw[j]
        if swaps:
            if key in visited:
                break
            visited[key] = True
        nb_ = sv.shape[0]
        lp = np.empty(nb_)
        ln_ = np.empty(nb_)
        pp = np.empty(nb_)
        pn = np.empty(nb_)
        wq = np.empty(nb_)
        T = bin_stats(sv, Wp, Wn, a, b, lp, ln_, pp, pn, wq)
        cntS = 0
        for j in range(d):
            if w[j] != 0.0:
                cntS += 1
        S = np.empty(cntS, np.int64)
        free = np.empty(d, np.bool_)
        u = 0
        nfree = 0
        for j in range(d):
            if w[j] != 0.0:
                S[u] = j
                u += 1
            free[j] = w[j] == 0.0 and valid[j]
            nfree += free[j]
        SHP, SHN = support_hist(S, inv, nb_, y, c, QT, var_of, kind, ncode, lev_of, SLIDER)
        oe, oaj, ov, ofx = main_phase(w, S, free, k, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of,
                                      lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp, allv, nzv,
                                      ne, tr_a, tr_b, add_screen, SHP, SHN, Wp, Wn)
        orj = np.full(oe.shape[0], -1, np.int64)
        sel = top_order(oe, m_chk)
        bi, L2, s2, inv2, sv2, Wp2, Wn2, a2, b2 = check_kernel(s, w, L, oe[sel], orj[sel], oaj[sel], ov[sel], a, b,
                                                               margin, y, c, QT, var_of, tptr, tabv, kind, lev_of,
                                                               inv, sv)
        rj = -1
        aj = -1
        v = 0.0
        from_main = False
        if bi >= 0:
            rj = orj[sel[bi]]
            aj = oaj[sel[bi]]
            v = ov[sel[bi]]
            from_main = True
            m_sel = sel
            m_bi = bi
        if bi < 0 and cntS > 0 and nfree > 0 and swaps and L / N <= gate:
            # threshold slides before the full swap phase
            se, srj, saj, sv_ = slide_moves(w, S, free, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of,
                                            lev_of, kind, ncode, vrp, vrows, vcols, vcp, ne, tr_a, tr_b,
                                            SLIDER, SHP, SHN)
            sel = top_order(se, m_chk)
            bi, L2, s2, inv2, sv2, Wp2, Wn2, a2, b2 = check_kernel(s, w, L, se[sel], srj[sel], saj[sel], sv_[sel],
                                                                   a, b, margin, y, c, QT, var_of, tptr, tabv, kind,
                                                                   lev_of, inv, sv)
            if bi >= 0:
                rj = srj[sel[bi]]
                aj = saj[sel[bi]]
                v = sv_[sel[bi]]
        if bi >= 0:
            pass
        elif nfree > 0 and swaps and L / N <= gate:
            oe, orj, oaj, ov, ofx = swap_phase(s, w, S, free, y, c, a, b, QT, var_of, lev_of, kind, ncode, tptr, tabv,
                                               vrp, vrows, vcols, vcp, nzv, n_screen, ne, tr_a, tr_b, inv, sv)
            sel = top_order(oe, m_chk)
            bi, L2, s2, inv2, sv2, Wp2, Wn2, a2, b2 = check_kernel(s, w, L, oe[sel], orj[sel], oaj[sel], ov[sel], a,
                                                                   b, margin, y, c, QT, var_of, tptr, tabv, kind,
                                                                   lev_of, inv, sv)
            if bi >= 0:
                rj = orj[sel[bi]]
                aj = oaj[sel[bi]]
                v = ov[sel[bi]]
            else:
                has_c = True
                c_oe, c_orj, c_oaj, c_ov = oe, orj, oaj, ov
        if bi < 0:
            break
        if rj >= 0:
            w[rj] = 0.0
        w[aj] = v
        s, L, a, b, inv, sv, Wp, Wn = s2, L2, a2, b2, inv2, sv2, Wp2, Wn2
        if from_main:
            # the other checked value changes (other columns), exactly on the new point, before a new main phase
            used = np.zeros(1, np.int64)
            used[0] = aj
            for t in range(m_sel.shape[0]):
                if t == m_bi:
                    continue
                q = m_sel[t]
                cj = oaj[q]
                if orj[q] >= 0 or w[cj] == 0.0 or cj == aj:
                    continue  # value changes of support columns only
                dup = False
                for u in range(used.shape[0]):
                    if used[u] == cj:
                        dup = True
                if dup:
                    continue
                one_e = np.full(1, -np.inf)
                bi3, L3, s3, inv3, sv3, Wp3, Wn3, a3, b3 = check_kernel(s, w, L, one_e, orj[q:q + 1], oaj[q:q + 1],
                                                                       ov[q:q + 1], a, b, margin, y, c, QT, var_of,
                                                                       tptr, tabv, kind, lev_of, inv, sv)
                if bi3 >= 0:
                    w[cj] = ov[q]
                    s, L, a, b, inv, sv, Wp, Wn = s3, L3, a3, b3, inv3, sv3, Wp3, Wn3
                    used = np.append(used, cj)
    return L / N, w, has_c, c_oe, c_orj, c_oaj, c_ov


@_kernel()
def start_state(w, QT, var_of, tptr, tabv, y, c, N):
    """ScoreState of the points w in one kernel: scores, bins and the calibrated (L, a, b)."""
    d = w.shape[0]
    cnt = 0
    for j in range(d):
        if w[j] != 0.0:
            cnt += 1
    S = np.empty(cnt, np.int64)
    wS = np.empty(cnt)
    u = 0
    for j in range(d):
        if w[j] != 0.0:
            S[u] = j
            wS[u] = w[j]
            u += 1
    s = score_kernel(QT, var_of, tptr, tabv, S, wS)
    inv, sv, Wp, Wn = bin_scores(s, y, c)
    if sv.shape[0] < 2:
        npos = 0.0
        for i in range(y.shape[0]):
            if y[i] > 0:
                npos += c[i]
        L, a, b = calibrate_bins(sv, Wp, Wn, 0.0, np.log(npos / (N - npos)))
    else:
        sd = np.std(s)
        L, a, b = calibrate_bins(sv / sd, Wp, Wn, 0.0, 0.0)
        a = a / sd
    return s, L, a, b, inv, sv, Wp, Wn


class ILS:
    """Best-improvement local search over integer points, scored by the calibrated loss."""

    def __init__(self, D, k, bound):
        self.D, self.k = D, k
        self.est_margin = 1e-4
        self.hw = _normals(11, D.d)  # hash of a point: its random projection
        self.vis = _new_dict_f64()
        self.cands = {}  # final point -> candidates of its failing swap phase
        self.allv = np.arange(-bound, bound + 1, dtype=np.float64)
        self.nzv = self.allv[self.allv != 0]

    def run(self, w, max_iter=100, gate=np.inf):
        D = self.D
        w = w.astype(np.float64).copy()
        st = start_state(w, D.QT, D.var_of, D.tptr, D.tabv, D.y, D.c, D.N)
        out = ils_kernel(w, st[0], st[1], st[2], st[3], st[4], st[5], st[6], st[7], self.vis, self.hw, self.k,
                         max_iter, True, float(gate), D.N, NEXACT, NSCREEN, ADDSCREEN, self.est_margin, D.y, D.c,
                         D.QT, D.var_of, D.lev_of, D.kind, D.ncode, D.tptr, D.tabv, D.vrp, D.vrows, D.vcols, D.vcp,
                         D.valid, self.allv, self.nzv, TR_A, TR_B)
        if out[2]:
            self.cands[out[1].tobytes()] = out[3:]
        return out[0], out[1]


# ------------------------------------------------------------------ model
class _Search:
    """One fit: beam search, calibrated rounding, local search, refit swaps and the exact polish."""

    def __init__(self, k, bound, time_limit, profile):
        self.k, self.bound, self.time_limit = k, bound, time_limit
        self.prof = profile
        self.stopped = False

    def _late(self, t0, frac):
        """True (and the fit is marked as stopped early) once frac of the time limit has passed."""
        if time.perf_counter() > t0 + frac * self.time_limit:
            self.stopped = True
            return True
        return False

    def _rswap(self, D, ils, w0, gate=np.inf, nsw=RSWAP):
        """Swaps with every point refitted: the nsw best swaps by the estimate at w0 (the candidates of its failing
        swap phase when its run kept them), each followed by a continuous refit of its support and a calibrated
        rounding; local search from the best rounding. Returns (loss, points) or None."""
        if np.count_nonzero(w0) < 2:
            return None
        key_b = w0.tobytes()
        if key_b in ils.cands:
            oe, orj, oaj, ov = ils.cands[key_b]
        else:
            oe, orj, oaj, ov, _ = _swaps_at(D, ils, w0)
        o = np.argsort(oe, kind="mergesort")[:nsw]
        bw, bl = None, np.inf
        for t in o:
            if not np.isfinite(oe[t]):
                break
            w1 = w0.copy()
            w1[orj[t]] = 0.0
            w1[oaj[t]] = ov[t]
            wr, lr = refit_round(w1, D.QT, D.var_of, D.lev_of, D.kind, D.ncode, D.tptr, D.tabv, D.y, D.c,
                                 float(self.bound), 1e-8, NMULT)
            if np.isfinite(lr) and lr < bl and wr.tobytes() not in self.rs_tried:
                bw, bl = wr, lr
        if bw is None:
            return None
        self.rs_tried.add(bw.tobytes())
        return ils.run(bw, gate=gate)

    def _polish(self, D, ils, best_l, best_w, t0):
        """Exact polish of the best point: every value change of a support column checked exactly, and swaps
        screened at removal states with a refitted score-to-risk map; a local search from any improvement. Only on
        near-separable fits, where the move estimate (one Newton step in (a, b)) is unreliable."""
        # the calibrated map spans a wide logit range ...
        stg = start_state(best_w, D.QT, D.var_of, D.tptr, D.tabv, D.y, D.c, D.N)
        if abs(stg[2]) * (stg[5].max() - stg[5].min()) < POLG:
            return best_l, best_w
        # ... and the loss is low (a wide logit range from heavy-tailed raw columns alone does not count)
        p1 = float(D.c[D.y > 0].sum()) / D.N
        if 0.0 < p1 < 1.0 and best_l >= -POLH * (p1 * np.log(p1) + (1 - p1) * np.log(1 - p1)):
            return best_l, best_w
        for _ in range(POLISH):
            if self._late(t0, 0.9):
                break
            w0 = best_w
            st0 = start_state(w0, D.QT, D.var_of, D.tptr, D.tabv, D.y, D.c, D.N)
            S0 = np.flatnonzero(w0).astype(np.int64)
            rj, aj, vv = [], [], []
            eL, ej, ev = _polish_values(D, ils, w0, S0, st0)
            if ej >= 0:
                rj.append(-1); aj.append(ej); vv.append(ev)
            free0 = (w0 == 0) & D.valid
            if free0.any():
                xL, xr, xa, xv = _polish_swaps(D, ils, w0, S0, st0, free0)
                if xa >= 0 and xL < st0[1] * (1 - 1e-9) and (ej < 0 or xL < eL):
                    rj, aj, vv = [xr], [xa], [xv]
            if not rj:
                break
            bi, L2 = _check_moves(D, w0, st0, rj, aj, vv)
            if bi < 0:
                break
            w1 = w0.copy()
            if rj[bi] >= 0:
                w1[rj[bi]] = 0.0
            w1[aj[bi]] = vv[bi]
            l, w = ils.run(w1)
            l1 = L2 / D.N
            if l1 < l:
                l, w = l1, w1
            if l < best_l - 1e-12:
                best_l, best_w = l, w
            else:
                break
        return best_l, best_w

    def fit(self, X, y):
        """Returns (points, calibrated mean log loss, seconds per stage: data, beam, rounding, local search)."""
        prof, k, bound = self.prof, self.k, self.bound
        tm = [time.perf_counter()]
        t0 = tm[0]
        D = Data(X, y, prof["ms_sqrt"])
        tm.append(time.perf_counter())
        self.stopped, st = beam_search(D, k, float(bound), prof["parent_size"], prof["final_pool"], prof["screen"],
                                       prof["lastscreen"], prof["sigmax"], CHILD, t0 + 0.4 * self.time_limit)
        tm.append(time.perf_counter())
        seen = set()
        starts = []
        best_w, best_l = np.zeros(D.d), np.inf
        hv = _normals(7, D.n)
        RR, LL, KK = round_all(st[2], st[1], st[3], st[4], st[5], st[7], st[8], st[9], D.QT, D.var_of, D.tptr, D.tabv,
                               NMULT, float(bound), hv, D.N)
        for t in range(len(LL)):
            w = np.zeros(D.d)
            w[st[7][t, :st[8][t]]] = RR[t, :st[8][t]]
            l = float(LL[t])
            key = round(float(KK[t]), 6)  # equivalent points (duplicate columns) give one start
            if key in seen or not np.isfinite(l):
                continue
            seen.add(key)
            starts.append((l, w))
            if l < best_l:
                best_l, best_w = l, w
        tm.append(time.perf_counter())
        # integer local search from the best few distinct rounded solutions
        ils = ILS(D, k, bound)
        self.rs_tried = set()
        starts.sort(key=lambda t: t[0])
        nofail = 0
        ends = []
        for l0, w0 in starts[:NSTARTS]:
            if self._late(t0, 0.8) or nofail >= SPAT or best_l < SEPEPS:
                break
            l, w = ils.run(w0)
            nofail = nofail + 1 if l >= best_l * (1.0 - SPEPS) - 1e-12 else 0
            ends.append((l, w))
            if l < best_l:
                best_l, best_w = l, w
        done = best_l < SEPEPS  # separable: no move can lower the loss by more than SEPEPS
        if not done and np.isfinite(best_l) and np.count_nonzero(best_w) > 1:
            # swaps with every point refitted, from the best point (RSF fails in all)
            fails = 0
            while fails < RSF and not self._late(t0, 0.8):
                res = self._rswap(D, ils, best_w, gate=best_l)
                if res is not None and res[0] < best_l - 1e-12:
                    best_l, best_w = res
                    fails = 0
                else:
                    fails += 1
                if res is None:
                    break
        if not done:
            # refit swaps also from the next best distinct start ends (the final point depends less on which end won)
            ends.sort(key=lambda t: t[0])
            seen_e = {best_w.tobytes(), ends[0][1].tobytes()} if ends else {best_w.tobytes()}
            cnt = 1
            for le, we in ends[1:]:
                if cnt >= RSEND or self._late(t0, 0.8):
                    break
                if we.tobytes() in seen_e or np.count_nonzero(we) < 2 or le > best_l * (1 + RSEGAP):
                    continue
                seen_e.add(we.tobytes())
                cnt += 1
                res = self._rswap(D, ils, we, nsw=RSEN)
                if res is not None and res[0] < best_l - 1e-12:
                    best_l, best_w = res
        if not done and np.isfinite(best_l) and np.count_nonzero(best_w) > 0:
            best_l, best_w = self._polish(D, ils, best_l, best_w, t0)
        tm.append(time.perf_counter())
        if tm[-1] - t0 > 0.9 * self.time_limit:
            self.stopped = True
        return np.clip(np.round(best_w), -bound, bound), float(best_l), np.diff(tm)


def _polish_values(D, ils, w0, S0, st0):
    """(loss, column, value) of the best exact value change of a support column (column -1 if none improves)."""
    return exact_values(w0, S0, st0[0], st0[1], st0[2], st0[3], st0[4], st0[5], st0[6], st0[7], D.y, D.c, D.QT,
                        D.var_of, D.kind, D.ncode, D.lev_of, D.tptr, D.tabv, ils.allv)


def _polish_swaps(D, ils, w0, S0, st0, free0):
    """(loss, removed column or -1, added column, value) of the best exact swap / addition (added -1 if none)."""
    return exact_moves(st0[0], w0, S0, free0, ils.k, D.y, D.c, st0[2], st0[3], D.QT, D.var_of, D.lev_of, D.kind,
                       D.ncode, D.tptr, D.tabv, D.vrp, D.vrows, D.vcols, D.vcp, ils.nzv, POLNS, st0[1] * (1 - 1e-9))


def _check_moves(D, w0, st0, rj, aj, vv):
    """Index of the first move (rj[t] removed, aj[t] set to vv[t]) whose exact calibrated loss improves on w0's
    (-1 if none), and that loss."""
    return check_kernel(st0[0], w0, st0[1], np.full(len(rj), -np.inf), np.array(rj, np.int64),
                        np.array(aj, np.int64), np.array(vv, np.float64), st0[2], st0[3], 1e-4, D.y, D.c, D.QT,
                        D.var_of, D.tptr, D.tabv, D.kind, D.lev_of, st0[4], st0[5])[:2]


def _swaps_at(D, ils, w0):
    """The swap phase's candidates at the points w0 (estimate, removed column, added column, value, fixed-map loss)."""
    st0 = start_state(w0, D.QT, D.var_of, D.tptr, D.tabv, D.y, D.c, D.N)
    S0 = np.flatnonzero(w0).astype(np.int64)
    free0 = (w0 == 0) & D.valid
    return swap_phase(st0[0], w0, S0, free0, D.y, D.c, st0[2], st0[3], D.QT, D.var_of, D.lev_of, D.kind, D.ncode,
                      D.tptr, D.tabv, D.vrp, D.vrows, D.vcols, D.vcp, ils.nzv, NSCREEN, NEXACT, TR_A, TR_B, st0[4],
                      st0[5])


_WARM = []


def _warmup():
    """Compile (or load from numba's disk cache) every kernel signature a fit uses, once per process, by fits of both
    profiles on small binary, small-integer and continuous data, so that no compilation happens inside a timed fit."""
    if _WARM:
        return
    with _LOCK:
        if not _WARM:  # another thread may have warmed up while this one waited
            _warm_fits()


def _warm_fits():
    _jit()
    rng = np.random.default_rng(0)
    n = 200
    X = np.hstack([(rng.random((n, 20)) < 0.4).astype(float), rng.integers(1, 6, size=(n, 2)).astype(float),
                   rng.random((n, 12)) * 10])
    y = (X[:, 0] + X[:, 4] / 5 + X[:, 6] / 10 + rng.random(n) > 1.5).astype(np.int64)
    # binary columns only: the byte code matrix from the start
    Xb = (rng.random((n, 16)) < 0.4).astype(float)
    yb = (Xb[:, 0] + Xb[:, 1] + rng.random(n) > 1.2).astype(np.int64)
    # more than 256 codes in a column: the int32 code matrix
    n = 400
    Xl = np.hstack([(rng.random((n, 10)) < 0.4).astype(float), rng.random((n, 2))])
    yl = (Xl[:, 0] + Xl[:, 10] + rng.random(n) > 1.2).astype(np.int64)
    for Xw, yw in ((X, y), (Xb, yb), (Xl, yl)):
        for prof in PROFILES.values():
            w0 = np.asarray(_Search(3, COEF_BOUND, 1e9, prof).fit(Xw, yw)[0], np.float64)
            # kernels called from Python only on some paths: the refit swaps' swap phase (when no failing swap phase
            # was kept) and the exact polish (near-separable fits only)
            D = Data(Xw, yw, prof["ms_sqrt"])
            ils = ILS(D, 3, COEF_BOUND)
            _swaps_at(D, ils, w0)
            st0 = start_state(w0, D.QT, D.var_of, D.tptr, D.tabv, D.y, D.c, D.N)
            S0 = np.flatnonzero(w0).astype(np.int64)
            _polish_values(D, ils, w0, S0, st0)
            _polish_swaps(D, ils, w0, S0, st0, (w0 == 0) & D.valid)
            _check_moves(D, w0, st0, [-1], [int(np.flatnonzero(w0 == 0)[0])], [1.0])
    _WARM.append(True)


def solve(X, y, k, bound=COEF_BOUND, time_limit=60.0, profile="decile"):
    """Integer points for the columns of ``X`` (y in {0, 1}): at most ``k`` nonzero, each in [-bound, bound],
    minimising the calibrated log loss. ``profile`` is "decile" (about 9 thresholds per numeric column) or "fine"
    (about 99). Returns (points, calibrated mean log loss, seconds per stage: data, beam, rounding, local search,
    stopped_early)."""
    if not HAVE_NUMBA:
        raise ImportError("the risk score solver needs numba (pip install numba)")
    if profile not in PROFILES:
        raise ValueError(f"profile must be one of {sorted(PROFILES)}, got {profile!r}")
    _warmup()
    search = _Search(int(k), int(bound), float(time_limit), PROFILES[profile])
    points, loss, seconds = search.fit(np.asarray(X, dtype=np.float64), np.asarray(y))
    return points, loss, seconds, search.stopped

Global variables

var HAVE_NUMBA

numba is optional for importing imodels but required to fit this model

var NUMBA_CACHE

compiled kernels are cached on disk (about 2 minutes to compile once per machine); set RISKSCORE_NUMBA_CACHE=0 to disable, e.g. when the package directory is read-only

var PROFILES

settings that differ between the two profiles (everything else is a module constant)

Functions

def beam_level(PINV, PNG, GOFF, GY, GC, GREP, GMG, PS, plen, PW, ploss, phash, psig, QT, var_of, lev_of, kind, ncode, vcp, vcols, tptr, tabv, vrp, vrows, bvalid, y, c, yc, allbin, colh, varh, pvar, seen, V, child_size, nfit, keep, sigmax, bound, tol)

One beam level on parents stored as arrays (row groups concatenated, offsets GOFF): Newton-gain screen of every column for every parent, proposals, exact child fits, sort by loss, at most sigmax children per signature, the keep best materialised as the next parents. Returns ok = False when nothing is proposed.

Expand source code
@_kernel()
def beam_level(PINV, PNG, GOFF, GY, GC, GREP, GMG, PS, plen, PW, ploss, phash, psig, QT, var_of, lev_of, kind,
               ncode, vcp, vcols, tptr, tabv, vrp, vrows, bvalid, y, c, yc, allbin, colh, varh, pvar, seen, V,
               child_size, nfit, keep, sigmax, bound, tol):
    """One beam level on parents stored as arrays (row groups concatenated, offsets GOFF): Newton-gain screen of
    every column for every parent, proposals, exact child fits, sort by loss, at most sigmax children per
    signature, the keep best materialised as the next parents. Returns ok = False when nothing is proposed."""
    P = PNG.shape[0]
    n = PINV.shape[1]
    d = var_of.shape[0]
    Pr = np.empty((n, P))
    for q in range(P):
        off = GOFF[q]
        sg = np.empty(PNG[q])
        for g in range(PNG[q]):
            sg[g] = _sig(-GMG[off + g])
        for i in range(n):
            Pr[i, q] = sg[PINV[q, i]]
    R, h00 = beam_rows_pr2(Pr, yc, c, not allbin)
    if allbin:  # x^2 = x: the third block equals the second
        G2 = colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, np.zeros(2 * P, np.bool_), vrp, vrows)
        Gm = np.empty((d, 3 * P))
        Gm[:, :2 * P] = G2
        Gm[:, 2 * P:] = G2[:, P:]
    else:
        pw = np.zeros(3 * P, np.bool_)
        pw[2 * P:] = True
        Gm = colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pw, vrp, vrows)
    GA, BA = newton_gains(Gm, h00, bvalid, bound)
    eq, ej, eest, eb, eh, es = beam_proposals(GA, BA, ploss, PS, plen, phash, psig, colh, varh, pvar, V, kind,
                                              vcols, vcp, seen, child_size, nfit, sigmax)
    m = eq.shape[0]
    lmax = PS.shape[1]
    if m == 0:
        return (False, PINV, PNG, GOFF, GY, GC, GREP, GMG, PS, plen, PW, ploss, phash, psig)
    # exact child fits, in the order: parents ascending; per parent its binary / small-code columns sorted by
    # variable (stable), then the others
    closs = np.empty(m)
    cq = np.empty(m, np.int64)
    cj = np.empty(m, np.int64)
    ch_ = np.empty(m, np.uint64)
    cs_ = np.empty(m, np.uint64)
    CW = np.zeros((m, lmax + 2))
    CS = np.zeros((m, lmax + 1), np.int64)
    u = 0
    for q in range(P):
        cnt = 0
        for t in range(m):
            if eq[t] == q:
                cnt += 1
        if cnt == 0:
            continue
        lst = np.empty(cnt, np.int64)
        cnt = 0
        for t in range(m):
            if eq[t] == q:
                lst[cnt] = t
                cnt += 1
        off = GOFF[q]
        ng = PNG[q]
        par_S = PS[q, :plen[q]].copy()
        par_w = PW[q, :plen[q] + 1].copy()
        gy = GY[off:off + ng]
        gc = GC[off:off + ng]
        rep = GREP[off:off + ng]
        nsp = 0
        for t in range(cnt):
            if kind[var_of[ej[lst[t]]]] < 2:
                nsp += 1
        if nsp > 0:
            idx = np.empty(nsp, np.int64)
            vv = np.empty(nsp, np.int64)
            nsp = 0
            for t in range(cnt):
                if kind[var_of[ej[lst[t]]]] < 2:
                    idx[nsp] = lst[t]
                    vv[nsp] = var_of[ej[lst[t]]]
                    nsp += 1
            idx = idx[np.argsort(vv, kind="mergesort")]
            js = ej[idx].copy()
            b0s = eb[idx].copy()
            losses, W = child_fit_batch(PINV[q], ng, gy, gc, rep, par_S, par_w, js, QT, var_of, lev_of, kind, ncode,
                                        tptr, tabv, c, bound, tol, False, vrp, vrows, b0s)
            for t in range(nsp):
                closs[u] = losses[t]
                cq[u] = q
                cj[u] = js[t]
                ch_[u] = eh[idx[t]]
                cs_[u] = es[idx[t]]
                CW[u, :plen[q] + 2] = W[t]
                u += 1
        for t in range(cnt):
            tt = lst[t]
            j = ej[tt]
            if kind[var_of[j]] < 2:
                continue
            v = var_of[j]
            code = np.empty(n, np.int64)
            for i in range(n):
                code[i] = QT[v, i]
            S, inv_, ng_, gy_, gc_, rep_, loss, mg_, w = make_child(PINV[q], ng, par_S, par_w, j, code, ncode[v], y,
                                                                    c, QT, var_of, tptr, tabv, bound, tol)
            closs[u] = loss
            cq[u] = q
            cj[u] = j
            ch_[u] = eh[tt]
            cs_[u] = es[tt]
            CW[u, :plen[q] + 2] = w
            u += 1
    # supports of the children (sorted)
    for t in range(m):
        q = cq[t]
        ps = plen[q]
        j = cj[t]
        r = 0
        ins = False
        for z in range(ps + 1):
            if not ins and (r >= ps or j < PS[q, r]):
                CS[t, z] = j
                ins = True
            else:
                CS[t, z] = PS[q, r]
                r += 1
    # a child with a coefficient that rounds to 0 at every scale of the rounding grid is not representable by
    # integer points (its support collapses to a smaller one after rounding): rank it after the others, when that
    # column is non-binary (e.g. a raw column on a large scale)
    for t in range(m):
        ps = plen[cq[t]] + 1
        top = 0.0
        for r in range(1, ps + 1):
            top = max(top, abs(CW[t, r]))
        for r in range(1, ps + 1):
            if abs(CW[t, r]) < REPR * top and kind[var_of[CS[t, r - 1]]] != 0:
                closs[t] += REPP
                break
    o = np.argsort(closs, kind="mergesort")
    sel = np.empty(min(m, keep), np.int64)
    ns = 0
    if sigmax > 0:
        sc = Dict.empty(key_type=nb.types.uint64, value_type=nb.types.int64)
        for t in range(m):
            if ns >= keep:
                break
            i = o[t]
            c0 = sc.get(cs_[i], 0)
            if c0 < sigmax:
                sc[cs_[i]] = c0 + 1
                sel[ns] = i
                ns += 1
    else:
        for t in range(min(m, keep)):
            sel[t] = o[t]
        ns = min(m, keep)
    sel = sel[:ns]
    # materialise the kept children as the next parents
    P2 = ns
    PINV2 = np.empty((P2, n), np.int64)
    PNG2 = np.empty(P2, np.int64)
    GOFF2 = np.zeros(P2 + 1, np.int64)
    PS2 = np.zeros((P2, lmax + 1), np.int64)
    plen2 = np.empty(P2, np.int64)
    PW2 = np.zeros((P2, lmax + 2))
    ploss2 = np.empty(P2)
    phash2 = np.empty(P2, np.uint64)
    psig2 = np.empty(P2, np.uint64)
    parts = List()
    for t in range(P2):
        i = sel[t]
        q = cq[i]
        ps = plen[q] + 1
        S = CS[i, :ps].copy()
        w = CW[i, :ps + 1].copy()
        inv_, ng_, gy_, gc_, rep_, mg_ = materialise_kernel(PINV[q], PNG[q], cj[i], S, w, QT, var_of, kind, lev_of,
                                                            ncode, tptr, tabv, y, c)
        PINV2[t] = inv_
        PNG2[t] = ng_
        GOFF2[t + 1] = GOFF2[t] + ng_
        PS2[t, :ps] = S
        plen2[t] = ps
        PW2[t, :ps + 1] = w
        ploss2[t] = closs[i]
        phash2[t] = ch_[i]
        psig2[t] = cs_[i]
        parts.append((gy_, gc_, rep_, mg_))
    tot = GOFF2[P2]
    GY2 = np.empty(tot)
    GC2 = np.empty(tot)
    GREP2 = np.empty(tot, np.int64)
    GMG2 = np.empty(tot)
    for t in range(P2):
        gy_, gc_, rep_, mg_ = parts[t]
        a0 = GOFF2[t]
        GY2[a0:a0 + PNG2[t]] = gy_
        GC2[a0:a0 + PNG2[t]] = gc_
        GREP2[a0:a0 + PNG2[t]] = rep_
        GMG2[a0:a0 + PNG2[t]] = mg_
    return (True, PINV2, PNG2, GOFF2, GY2, GC2, GREP2, GMG2, PS2, plen2, PW2, ploss2, phash2, psig2)
def beam_levels(st, lev0, nlev, QT, var_of, lev_of, kind, ncode, vcp, vcols, tptr, tabv, vrp, vrows, bvalid, y, c, yc, allbin, colh, varh, pvar, seen, V, child_size, parent_size, final_pool, screen, lastscreen, sigmax, bound, tol)

Beam levels lev0 .. nlev - 1 in one kernel (same steps as the per-level loop of beam_search).

Expand source code
@_kernel()
def beam_levels(st, lev0, nlev, QT, var_of, lev_of, kind, ncode, vcp, vcols, tptr, tabv, vrp, vrows, bvalid, y,
                c, yc, allbin, colh, varh, pvar, seen, V, child_size, parent_size, final_pool, screen, lastscreen, sigmax,
                bound, tol):
    """Beam levels lev0 .. nlev - 1 in one kernel (same steps as the per-level loop of beam_search)."""
    for lev in range(lev0, nlev):
        last = lev == nlev - 1
        keep = final_pool if last else parent_size
        nfit = int((lastscreen if last else screen) * keep)
        w_ = max(lev, 1)
        res = beam_level(st[0], st[1], st[2], st[3], st[4], st[5], st[6], np.ascontiguousarray(st[7][:, :w_]),
                         st[8], np.ascontiguousarray(st[9][:, :w_ + 1]), st[10], st[11], st[12], QT, var_of, lev_of,
                         kind, ncode, vcp, vcols, tptr, tabv, vrp, vrows, bvalid, y, c, yc, allbin, colh, varh,
                         pvar, seen, V, child_size, nfit, keep, sigmax, bound, tol)
        if not res[0]:
            break
        st = (res[1], res[2], res[3], res[4], res[5], res[6], res[7], res[8], res[9], res[10], res[11], res[12],
              res[13])
    return st
def beam_proposals(GA, BA, ploss, PS, plen, phash, psig, colh, varh, var_of, V, kind, vcols, vcp, seen, child_size, nfit, sigmax)

Per parent its child_size best variables (each at its best column by the Newton gain), new supports only (seen: hashes of supports already proposed); sorted by estimated loss, at most sigmax per multiset of variables, the nfit best. Returns (parent, column, estimated loss, warm start, support hash, signature).

Expand source code
@_kernel()
def beam_proposals(GA, BA, ploss, PS, plen, phash, psig, colh, varh, var_of, V, kind, vcols, vcp, seen,
                   child_size, nfit, sigmax):
    """Per parent its child_size best variables (each at its best column by the Newton gain), new supports only
    (seen: hashes of supports already proposed); sorted by estimated loss, at most sigmax per multiset of
    variables, the nfit best. Returns (parent, column, estimated loss, warm start, support hash, signature)."""
    P, d = GA.shape
    maxp = P * child_size
    e_est = np.empty(maxp)
    e_q = np.empty(maxp, np.int64)
    e_j = np.empty(maxp, np.int64)
    e_b = np.empty(maxp)
    e_h = np.empty(maxp, np.uint64)
    e_s = np.empty(maxp, np.uint64)
    u = 0
    for q in range(P):
        gain = GA[q].copy()
        Sq = PS[q, :plen[q]]
        for t in range(plen[q]):
            gain[Sq[t]] = -1.0
        pick = pick_per_var(gain, var_of, V, child_size)  # at most one column per threshold cell
        for t in range(pick.shape[0]):
            j = pick[t]
            key = phash[q] + colh[j]
            if key in seen:
                continue
            seen[key] = True
            e_est[u] = ploss[q] - gain[j]
            e_q[u] = q
            e_j[u] = j
            e_b[u] = BA[q, j]
            e_h[u] = key
            e_s[u] = psig[q] + varh[var_of[j]]
            u += 1
    o = np.argsort(e_est[:u], kind="mergesort")
    cnt = Dict.empty(key_type=nb.types.uint64, value_type=nb.types.int64)
    sel = np.empty(min(u, nfit), np.int64)
    m = 0
    for t in range(u):
        if m >= nfit:
            break
        i = o[t]
        c0 = cnt.get(e_s[i], 0)
        if c0 < sigmax:
            cnt[e_s[i]] = c0 + 1
            sel[m] = i
            m += 1
    sel = sel[:m]
    return e_q[sel], e_j[sel], e_est[sel], e_b[sel], e_h[sel], e_s[sel]
def beam_rows_pr2(Pr, yc, c, third)

beam_rows from sigma(-m) per row; the x^2 block only when third; also the column sums of the curvature.

Expand source code
@_kernel()
def beam_rows_pr2(Pr, yc, c, third):
    """beam_rows from sigma(-m) per row; the x^2 block only when `third`; also the column sums of the curvature."""
    n, P = Pr.shape
    R = np.empty((n, (3 if third else 2) * P))
    h00 = np.zeros(P)
    for i in range(n):
        for q in range(P):
            pr = Pr[i, q]
            R[i, q] = yc[i] * pr
            h = c[i] * pr * (1.0 - pr)
            R[i, P + q] = h
            h00[q] += h
            if third:
                R[i, 2 * P + q] = h
    return R, h00

Beam search over supports, one numba kernel per level (beam_level). Every column's gain for every parent is bounded by one Newton step on the new coefficient (from per-code histograms of the gradient and curvature rows); each parent proposes its child_size best variables (each at its best column), and only the best screen x parent_size proposals (lastscreen x final_pool at the last level) are fitted exactly; at most sigmax children per multiset of variables are kept. Returns the final beam and whether it ran out of time.

Expand source code
def beam_search(D, k, bound, parent_size, final_pool, screen, lastscreen, sigmax, child_size=CHILD,
                deadline=np.inf):
    """Beam search over supports, one numba kernel per level (beam_level). Every column's gain for every parent is
    bounded by one Newton step on the new coefficient (from per-code histograms of the gradient and curvature
    rows); each parent proposes its child_size best variables (each at its best column), and only the best
    screen x parent_size proposals (lastscreen x final_pool at the last level) are fitted exactly; at most sigmax
    children per multiset of variables are kept. Returns the final beam and whether it ran out of time."""
    pvar, PV = D.pvar, D.npvar  # the variable (or the threshold cell of a variable) each column belongs to
    hi5 = _ints5(D.d + max(D.V, PV))
    colh = hi5[:D.d].astype(np.uint64) * np.uint64(2) + np.uint64(1)
    colh = colh[D.crep]  # a complement pair has one hash: mirrored supports are proposed once
    varh = hi5[D.d:D.d + PV].astype(np.uint64) * np.uint64(2) + np.uint64(1)
    inv, ng, gy, gc, rep = regroup(np.zeros(D.n, np.int64), 1, (D.y > 0).astype(np.int64), 2, D.y, D.c)
    npos = D.c[D.y > 0].sum()
    w0 = np.log(npos / (D.N - npos))
    mg = gy * w0
    lmax = max(k, 1)
    st = (inv[None, :].copy(), np.array([ng], np.int64), np.array([0, ng], np.int64), gy, gc, rep, mg,
          np.zeros((1, lmax), np.int64), np.zeros(1, np.int64), np.full((1, lmax + 1), w0),
          np.array([total_loss(mg, gc)]), np.zeros(1, np.uint64), np.zeros(1, np.uint64))
    seen = _new_dict_u64()
    nlev = min(k, int(D.valid.sum()))
    t_start = time.perf_counter()
    stopped = False
    for lev in range(nlev):
        if lev == 1:
            # the remaining levels in one kernel when they surely fit in the time left (level 0, one parent, took
            # dt; a level of the full beam costs at most ~parent_size times that)
            dt = time.perf_counter() - t_start
            if t_start + 2.0 * parent_size * nlev * max(dt, 1e-3) < deadline:
                return False, beam_levels(st, 1, nlev, D.QT, D.var_of, D.lev_of, D.kind, D.ncode, D.vcp, D.vcols,
                                          D.tptr, D.tabv, D.vrp, D.vrows, D.bvalid, D.y, D.c, D.yc, D.allbin, colh,
                                          varh, pvar, seen, PV, child_size, parent_size, final_pool, screen,
                                          lastscreen, sigmax, bound, BTOL)
        last = lev == nlev - 1
        if time.perf_counter() > deadline and st[1].shape[0] > 1:
            # out of time: finish the support greedily from the best parent
            stopped = True
            ng0 = st[1][0]
            st = (st[0][:1], st[1][:1], st[2][:2], st[3][:ng0], st[4][:ng0], st[5][:ng0], st[6][:ng0], st[7][:1],
                  st[8][:1], st[9][:1], st[10][:1], st[11][:1], st[12][:1])
        keep = final_pool if last else parent_size
        nfit = int((lastscreen if last else screen) * keep)
        res = beam_level(st[0], st[1], st[2], st[3], st[4], st[5], st[6], np.ascontiguousarray(st[7][:, :max(lev, 1)]),
                         st[8], np.ascontiguousarray(st[9][:, :max(lev, 1) + 1]), st[10], st[11], st[12], D.QT,
                         D.var_of, D.lev_of, D.kind, D.ncode, D.vcp, D.vcols, D.tptr, D.tabv, D.vrp, D.vrows,
                         D.bvalid, D.y, D.c, D.yc, D.allbin, colh, varh, pvar, seen, PV, child_size, nfit, keep, sigmax,
                         bound, BTOL)
        if not res[0]:
            break
        st = res[1:]
    return stopped, st
def bin_scores(s, y, c)

Group rows by distinct score: inv (row -> bin), bin scores sv, weights of y=+1 (Wp) and y=-1 (Wn).

Expand source code
@_kernel()
def bin_scores(s, y, c):
    """Group rows by distinct score: inv (row -> bin), bin scores sv, weights of y=+1 (Wp) and y=-1 (Wn)."""
    n = s.shape[0]
    lo = s[0]
    hi = s[0]
    integral = True
    for i in range(n):
        v = s[i]
        if v < lo:
            lo = v
        if v > hi:
            hi = v
        if integral and v != np.floor(v):
            integral = False
    if integral and hi - lo <= 4 * n + 1024:
        # integer scores in a small range: counting instead of sorting
        R = int(hi - lo) + 1
        cid = np.full(R, -1, np.int64)
        for i in range(n):
            cid[int(s[i] - lo)] = 0
        nb_ = 0
        for r in range(R):
            if cid[r] >= 0:
                cid[r] = nb_
                nb_ += 1
        inv = np.empty(n, np.int64)
        sv = np.empty(nb_)
        Wp = np.zeros(nb_)
        Wn = np.zeros(nb_)
        for r in range(R):
            if cid[r] >= 0:
                sv[cid[r]] = lo + r
        for i in range(n):
            q = cid[int(s[i] - lo)]
            inv[i] = q
            if y[i] > 0:
                Wp[q] += c[i]
            else:
                Wn[q] += c[i]
        return inv, sv, Wp, Wn
    order = np.argsort(s)
    inv = np.empty(n, np.int64)
    sv = np.empty(n)
    Wp = np.zeros(n)
    Wn = np.zeros(n)
    nb_ = -1
    last = 0.0
    for t in range(n):
        i = order[t]
        if t == 0 or s[i] != last:
            nb_ += 1
            sv[nb_] = s[i]
            last = s[i]
        inv[i] = nb_
        if y[i] > 0:
            Wp[nb_] += c[i]
        else:
            Wn[nb_] += c[i]
    nb_ += 1
    return inv, sv[:nb_].copy(), Wp[:nb_].copy(), Wn[:nb_].copy()
def bin_stats(sv, Wp, Wn, a, b, lp, ln_, pp, pn, wq)
Expand source code
@_kernel()
def bin_stats(sv, Wp, Wn, a, b, lp, ln_, pp, pn, wq):
    T = np.zeros(6)
    for q in range(sv.shape[0]):
        z = a * sv[q] + b
        e = np.exp(-abs(z))
        l1 = np.log1p(e)
        if z > 0:
            lp[q] = l1
            ln_[q] = z + l1
            pp[q] = e / (1.0 + e)
            pn[q] = 1.0 / (1.0 + e)
        else:
            lp[q] = -z + l1
            ln_[q] = l1
            pp[q] = 1.0 / (1.0 + e)
            pn[q] = e / (1.0 + e)
        wq[q] = pp[q] * pn[q]
        g = -Wp[q] * pp[q] + Wn[q] * pn[q]
        wt = (Wp[q] + Wn[q]) * wq[q]
        T[0] += Wp[q] * lp[q] + Wn[q] * ln_[q]
        T[1] += g * sv[q]
        T[2] += g
        T[3] += wt * sv[q] * sv[q]
        T[4] += wt * sv[q]
        T[5] += wt
    return T
def build_chains(B, cnt, order, d)

Greedy chain cover of the columns in order (by count descending): each column joins the chain whose last column contains it with the smallest count, else starts a new chain. Returns chain id and level (1-based position) of every column (-1 when not in order), and the number of chains.

Expand source code
@_kernel()
def build_chains(B, cnt, order, d):
    """Greedy chain cover of the columns in `order` (by count descending): each column joins the chain whose
    last column contains it with the smallest count, else starts a new chain. Returns chain id and level
    (1-based position) of every column (-1 when not in `order`), and the number of chains."""
    nw = B.shape[0]
    m = order.shape[0]
    last = np.empty(m, np.int64)
    length = np.zeros(m, np.int64)
    chain = np.full(d, -1, np.int64)
    lev = np.zeros(d, np.int64)
    nch = 0
    for t in range(m):
        j = order[t]
        best = -1
        for ch in range(nch):
            l = last[ch]
            if cnt[l] < cnt[j]:
                continue
            if best >= 0 and cnt[l] >= cnt[last[best]]:
                continue
            ok = True
            for w in range(nw):
                if B[w, j] & ~B[w, l]:
                    ok = False
                    break
            if ok:
                best = ch
        if best < 0:
            best = nch
            nch += 1
        last[best] = j
        length[best] += 1
        chain[j] = best
        lev[j] = length[best]
    return chain, lev, nch
def calib_round_kernel(XS, gy, gc, beta, n_mult, bound)

Round m * beta for a grid of scales (largest point 0.5 .. bound + 0.49); return the rounding with the smallest calibrated loss on the row groups (XS: group rows of the support columns).

Expand source code
@_kernel()
def calib_round_kernel(XS, gy, gc, beta, n_mult, bound):
    """Round m * beta for a grid of scales (largest point 0.5 .. bound + 0.49); return the rounding with the
    smallest calibrated loss on the row groups (XS: group rows of the support columns)."""
    ng, p = XS.shape
    top = 0.0
    for q in range(p):
        top = max(top, abs(beta[q]))
    best_l = np.inf
    best_r = np.zeros(p)
    prev = np.full(p, np.nan)
    r = np.empty(p)
    sg = np.empty(ng)
    if top < 1e-12:
        return best_r, best_l
    seen_r = np.empty((n_mult + RCLIPN, p))
    nseen = 0
    WPc = np.zeros(0)
    WNc = np.zeros(0)
    pres = np.zeros(0, np.bool_)
    svb = np.empty(0)
    Wpb = np.empty(0)
    Wnb = np.empty(0)
    Eb = np.empty(0)
    Enb = np.empty(0)
    for t in range(n_mult + RCLIPN):
        if t < n_mult:
            L = 0.5 + (bound - 0.01) * t / max(n_mult - 1, 1)
        else:
            # clipped scales: the largest points saturate at the bound, the others get finer ratios
            L = bound * (1.0 + (RCLIP - 1.0) * (t - n_mult + 1) / RCLIPN)
        same = True
        anynz = False
        for q in range(p):
            v = np.round(beta[q] * L / top)
            v = min(max(v, -bound), bound)
            r[q] = v
            if v != prev[q]:
                same = False
            if v != 0:
                anynz = True
        if same or not anynz:
            continue
        prev[:] = r
        # a rounding proportional to an earlier one has the same calibrated loss: skip it
        gg = 0
        for q in range(p):
            gg = _gcd(gg, int(abs(r[q])))
        sgn = 0.0
        for q in range(p):
            if r[q] != 0.0:
                sgn = 1.0 if r[q] > 0 else -1.0
                break
        dup = False
        for t2 in range(nseen):
            eq_ = True
            for q in range(p):
                if seen_r[t2, q] != sgn * r[q] / gg:
                    eq_ = False
                    break
            if eq_:
                dup = True
                break
        if dup:
            continue
        for q in range(p):
            seen_r[nseen, q] = sgn * r[q] / gg
        nseen += 1
        mn = np.inf
        mx = -np.inf
        m1 = 0.0
        m2 = 0.0
        for g in range(ng):
            v = 0.0
            for q in range(p):
                v += XS[g, q] * r[q]
            sg[g] = v
            mn = min(mn, v)
            mx = max(mx, v)
            m1 += gc[g] * v
            m2 += gc[g] * v * v
        if mx == mn:
            continue
        tot = gc.sum()
        sd = np.sqrt(max(m2 / tot - (m1 / tot) ** 2, 1e-300))
        # integer scores: calibrate on the distinct score values (counted in reused buffers)
        integral = True
        for g in range(ng):
            if sg[g] != np.floor(sg[g]):
                integral = False
                break
        if integral and mx - mn <= 4 * ng + 1024:
            R = int(mx - mn) + 1
            if R > WPc.shape[0]:
                WPc = np.zeros(R)
                WNc = np.zeros(R)
                pres = np.zeros(R, np.bool_)
                svb = np.empty(R)
                Wpb = np.empty(R)
                Wnb = np.empty(R)
                Eb = np.empty(2 * R)
                Enb = np.empty(2 * R)
            for t2 in range(R):
                WPc[t2] = 0.0
                WNc[t2] = 0.0
                pres[t2] = False
            for g in range(ng):
                ix = int(sg[g] - mn)
                pres[ix] = True
                if gy[g] > 0:
                    WPc[ix] += gc[g]
                else:
                    WNc[ix] += gc[g]
            mb_ = 0
            for t2 in range(R):
                if pres[t2]:
                    svb[mb_] = (mn + t2) / sd
                    Wpb[mb_] = WPc[t2]
                    Wnb[mb_] = WNc[t2]
                    mb_ += 1
            loss, _, _ = calibrate_bins_buf(svb[:mb_], Wpb[:mb_], Wnb[:mb_], 0.0, 0.0, Eb, Enb)
        else:
            inv_, sv, Wp, Wn = bin_scores(sg, gy, gc)
            loss, _, _ = calibrate_bins(sv / sd, Wp, Wn, 0.0, 0.0)
        if loss < best_l:
            best_l = loss
            best_r[:] = r
    return best_r, best_l
def calib_target(sv, Wp, Wn, a, b, target, E, En)

calibrate_bins_buf that stops early: returns the loss once it is below target (an upper bound of the calibrated loss, so the move surely beats target), or np.inf once even twice the Newton decrement cannot reach target (rejected), else the converged loss.

Expand source code
@_kernel()
def calib_target(sv, Wp, Wn, a, b, target, E, En):
    """calibrate_bins_buf that stops early: returns the loss once it is below target (an upper bound of the
    calibrated loss, so the move surely beats target), or np.inf once even twice the Newton decrement cannot reach
    target (rejected), else the converged loss."""
    m = sv.shape[0]
    cur = 0.0
    for i in range(m):
        lv, E[i] = _lrow_e(a * sv[i] + b)
        cur += Wp[i] * lv
    for i in range(m):
        lv, E[m + i] = _lrow_e(-(a * sv[i] + b))
        cur += Wn[i] * lv
    if cur < target:
        return cur
    for _ in range(CALCAP):
        ga = gb = haa = hab = hbb = 0.0
        for i in range(m):
            z = a * sv[i] + b
            p = _sig_e(z, E[i])
            w = Wp[i] * p * (1.0 - p)
            gi = -Wp[i] * p
            ga += gi * sv[i]
            gb += gi
            haa += w * sv[i] * sv[i]
            hab += w * sv[i]
            hbb += w
        for i in range(m):
            z = -(a * sv[i] + b)
            p = _sig_e(z, E[m + i])
            w = Wn[i] * p * (1.0 - p)
            gi = Wn[i] * p
            ga += gi * sv[i]
            gb += gi
            haa += w * sv[i] * sv[i]
            hab += w * sv[i]
            hbb += w
        haa += 1e-12
        hbb += 1e-12
        det = haa * hbb - hab * hab
        if det <= 1e-18 * (haa * hbb + 1e-300):
            da = 0.0
            db = gb / hbb
        else:
            da = (hbb * ga - hab * gb) / det
            db = (haa * gb - hab * ga) / det
        dec = ga * da + gb * db
        if cur - REJ * dec > target:
            return np.inf
        t = 1.0
        new = cur
        while t > 1e-10:
            na, nb_ = a - t * da, b - t * db
            new = 0.0
            for i in range(m):
                lv, En[i] = _lrow_e(na * sv[i] + nb_)
                new += Wp[i] * lv
            for i in range(m):
                lv, En[m + i] = _lrow_e(-(na * sv[i] + nb_))
                new += Wn[i] * lv
            if new <= cur - 1e-4 * t * dec:
                break
            t *= 0.5
        if t <= 1e-10:
            break
        a, b = na, nb_
        E, En = En, E
        improv = cur - new
        cur = new
        if cur < target:
            return cur
        if improv < 1e-12 * (1.0 + cur):
            break
    return cur
def calibrate_bins(sv, Wp, Wn, a, b)

calibrate() on the bins (score sv, weight Wp of y = +1, Wn of y = -1) without building the 2m-row arrays: the same terms in the same order (all y = +1 terms, then all y = -1 terms), so the same result.

Expand source code
@_kernel()
def calibrate_bins(sv, Wp, Wn, a, b):
    """calibrate() on the bins (score sv, weight Wp of y = +1, Wn of y = -1) without building the 2m-row arrays:
    the same terms in the same order (all y = +1 terms, then all y = -1 terms), so the same result."""
    m = sv.shape[0]
    return calibrate_bins_buf(sv, Wp, Wn, a, b, np.empty(2 * m), np.empty(2 * m))
def calibrate_bins_buf(sv, Wp, Wn, a, b, E, En)

calibrate_bins with caller work buffers E, En (length >= 2m).

Expand source code
@_kernel()
def calibrate_bins_buf(sv, Wp, Wn, a, b, E, En):
    """calibrate_bins with caller work buffers E, En (length >= 2m)."""
    m = sv.shape[0]
    cur = 0.0
    for i in range(m):
        lv, E[i] = _lrow_e(a * sv[i] + b)
        cur += Wp[i] * lv
    for i in range(m):
        lv, E[m + i] = _lrow_e(-(a * sv[i] + b))
        cur += Wn[i] * lv
    for _ in range(100):
        ga = gb = haa = hab = hbb = 0.0
        for i in range(m):
            z = a * sv[i] + b
            p = _sig_e(z, E[i])
            w = Wp[i] * p * (1.0 - p)
            gi = -Wp[i] * p
            ga += gi * sv[i]
            gb += gi
            haa += w * sv[i] * sv[i]
            hab += w * sv[i]
            hbb += w
        for i in range(m):
            z = -(a * sv[i] + b)
            p = _sig_e(z, E[m + i])
            w = Wn[i] * p * (1.0 - p)
            gi = Wn[i] * p
            ga += gi * sv[i]
            gb += gi
            haa += w * sv[i] * sv[i]
            hab += w * sv[i]
            hbb += w
        haa += 1e-12
        hbb += 1e-12
        det = haa * hbb - hab * hab
        if det <= 1e-18 * (haa * hbb + 1e-300):
            da = 0.0
            db = gb / hbb
        else:
            da = (hbb * ga - hab * gb) / det
            db = (haa * gb - hab * ga) / det
        dec = ga * da + gb * db
        t = 1.0
        new = cur
        while t > 1e-10:
            na, nb_ = a - t * da, b - t * db
            new = 0.0
            for i in range(m):
                lv, En[i] = _lrow_e(na * sv[i] + nb_)
                new += Wp[i] * lv
            for i in range(m):
                lv, En[m + i] = _lrow_e(-(na * sv[i] + nb_))
                new += Wn[i] * lv
            if new <= cur - 1e-4 * t * dec:
                break
            t *= 0.5
        if t <= 1e-10:
            break
        a, b = na, nb_
        E, En = En, E
        improv = cur - new
        cur = new
        if improv < 1e-12 * (1.0 + cur):
            break
    return cur, a, b
def chain_codes(B, n, chain, lev, nch, Q)

Code of every row in every chain: the largest level whose column contains the row (binary search, the columns of a chain are nested).

Expand source code
@_kernel()
def chain_codes(B, n, chain, lev, nch, Q):
    """Code of every row in every chain: the largest level whose column contains the row (binary search, the
    columns of a chain are nested)."""
    d = chain.shape[0]
    length = np.zeros(nch, np.int64)
    for j in range(d):
        if chain[j] >= 0:
            length[chain[j]] = max(length[chain[j]], lev[j])
    cp = np.zeros(nch + 1, np.int64)
    for ch in range(nch):
        cp[ch + 1] = cp[ch] + length[ch]
    cols = np.empty(cp[nch], np.int64)
    for j in range(d):
        if chain[j] >= 0:
            cols[cp[chain[j]] + lev[j] - 1] = j
    nw = B.shape[0]
    for ch in range(nch):
        # rows in level l but not in level l + 1 have code l (levels are nested): visit the set bits of the
        # differences, so each row is written once
        m = length[ch]
        base = cp[ch]
        for w in range(nw):
            nxt = np.uint64(0)
            for l in range(m - 1, -1, -1):
                cur = B[w, cols[base + l]]
                x = cur & ~nxt
                nxt = cur
                while x != np.uint64(0):
                    low = x & (~x + np.uint64(1))
                    i = w * 64 + DEBRUIJN[(low * np.uint64(0x03F79D71B4CA8B09)) >> np.uint64(58)]
                    Q[ch, i] = l + 1
                    x ^= low
    return Q
def chain_tables(tptr, lev_of, var_of, nch)

Value tables of the chain columns: x = 1 for codes >= level (other columns filled by the caller).

Expand source code
@_kernel()
def chain_tables(tptr, lev_of, var_of, nch):
    """Value tables of the chain columns: x = 1 for codes >= level (other columns filled by the caller)."""
    tabv = np.zeros(tptr[-1])
    for j in range(lev_of.shape[0]):
        if var_of[j] < nch:
            for t in range(tptr[j] + lev_of[j], tptr[j + 1]):
                tabv[t] = 1.0
    return tabv
def check_kernel(s, w, L, est, rjs, ajs, vs, a, b, margin, y, c, QT, var_of, tptr, tabv, kind, lev_of, cinv, csv)

Exact calibrated loss of the candidate moves (in order); returns the index of the best improving one (-1 if none) and its score vector, bins and calibration. Moves on binary columns are scored from rows keyed by (score bin, x_removed, x_added) in one pass.

Expand source code
@_kernel()
def check_kernel(s, w, L, est, rjs, ajs, vs, a, b, margin, y, c, QT, var_of, tptr, tabv, kind, lev_of, cinv, csv):
    """Exact calibrated loss of the candidate moves (in order); returns the index of the best improving one (-1 if
    none) and its score vector, bins and calibration. Moves on binary columns are scored from rows keyed by
    (score bin, x_removed, x_added) in one pass."""
    n = s.shape[0]
    bi = -1
    bL = np.inf
    ba = a
    bb = b
    nb0 = csv.shape[0]
    Wp4 = np.zeros(4 * nb0)
    Wn4 = np.zeros(4 * nb0)
    sc = np.empty(4 * nb0)
    kp = np.empty(4 * nb0, np.int64)
    for t in range(est.shape[0]):
        if est[t] > L * (1.0 + margin):
            continue  # the estimate says the move does not help
        rj = rjs[t]
        aj = ajs[t]
        va = var_of[aj]
        if kind[va] == 0 and (rj < 0 or kind[var_of[rj]] == 0):
            la = lev_of[aj]
            dv = vs[t] - w[aj]
            vr = var_of[rj] if rj >= 0 else 0
            lr = lev_of[rj] if rj >= 0 else 0
            wr = w[rj] if rj >= 0 else 0.0
            Wp4[:] = 0.0
            Wn4[:] = 0.0
            for i in range(n):
                kk = 4 * cinv[i] + (1 if QT[va, i] >= la else 0)
                if rj >= 0 and QT[vr, i] >= lr:
                    kk += 2
                if y[i] > 0:
                    Wp4[kk] += c[i]
                else:
                    Wn4[kk] += c[i]
            m = 0
            for kk in range(4 * nb0):
                if Wp4[kk] > 0.0 or Wn4[kk] > 0.0:
                    kp[m] = kk
                    sc[m] = csv[kk >> 2] - wr * ((kk >> 1) & 1) + dv * (kk & 1)
                    m += 1
            o = np.argsort(sc[:m], kind="mergesort")
            sv = np.empty(m)
            Wp = np.zeros(m)
            Wn = np.zeros(m)
            g = -1
            for u in range(m):
                kk = kp[o[u]]
                if g < 0 or sc[o[u]] != sv[g]:
                    g += 1
                    sv[g] = sc[o[u]]
                Wp[g] += Wp4[kk]
                Wn[g] += Wn4[kk]
            g += 1
            if g < 2:
                continue
            L2, a2, b2 = calibrate_bins(sv[:g], Wp[:g], Wn[:g], a, b)
        else:
            s2 = s.copy()
            if rj >= 0:
                v = var_of[rj]
                base = tptr[rj]
                for i in range(n):
                    s2[i] -= w[rj] * tabv[base + QT[v, i]]
            dv = vs[t] - w[aj]
            base = tptr[aj]
            for i in range(n):
                s2[i] += dv * tabv[base + QT[va, i]]
            lo = s2[0]
            hi = s2[0]
            for i in range(n):
                lo = min(lo, s2[i])
                hi = max(hi, s2[i])
            if hi == lo:
                continue
            inv, sv, Wp, Wn = bin_scores(s2, y, c)
            L2, a2, b2 = calibrate_bins(sv, Wp, Wn, a, b)
        if L2 < L - 1e-9 * L and L2 < bL:
            bi, bL, ba, bb = t, L2, a2, b2
            break  # first improving candidate in estimate order
    if bi < 0:
        return bi, bL, s, np.zeros(0, np.int64), np.zeros(0), np.zeros(0), np.zeros(0), a, b
    # the winner's score vector and bins
    s2 = s.copy()
    rj = rjs[bi]
    if rj >= 0:
        v = var_of[rj]
        base = tptr[rj]
        for i in range(n):
            s2[i] -= w[rj] * tabv[base + QT[v, i]]
    aj = ajs[bi]
    dv = vs[bi] - w[aj]
    v = var_of[aj]
    base = tptr[aj]
    for i in range(n):
        s2[i] += dv * tabv[base + QT[v, i]]
    inv, sv, Wp, Wn = bin_scores(s2, y, c)
    return bi, bL, s2, inv, sv, Wp, Wn, ba, bb
def child_fit_batch(par_inv, par_ng, par_gy, par_gc, par_rep, par_S, par_w, js, QT, var_of, lev_of, kind, ncode, tptr, tabv, c, bound, tol, screen, vrp, vrows, init)

Fit the children (parent support + column j) for the columns js (sorted by variable) of one parent without regrouping rows: the child's cells are the parent's groups split by the value of x_j, with weights from a (group, code) histogram of j's variable (suffix sums over codes for a chain).

Expand source code
@_kernel()
def child_fit_batch(par_inv, par_ng, par_gy, par_gc, par_rep, par_S, par_w, js, QT, var_of, lev_of, kind, ncode,
                    tptr, tabv, c, bound, tol, screen, vrp, vrows, init):
    """Fit the children (parent support + column j) for the columns js (sorted by variable) of one parent
    without regrouping rows: the child's cells are the parent's groups split by the value of x_j, with weights
    from a (group, code) histogram of j's variable (suffix sums over codes for a chain)."""
    m = js.shape[0]
    ps = par_S.shape[0]
    losses = np.empty(m)
    Wout = np.empty((m, ps + 2))
    n = par_inv.shape[0]
    Zp = np.empty((par_ng, ps))  # y_g * x_g of the parent's support columns
    for g in range(par_ng):
        for r in range(ps):
            Zp[g, r] = par_gy[g] * xval(QT, var_of, tptr, tabv, par_rep[g], par_S[r])
    offg = np.zeros(par_ng)  # margin of the parent's support part per group (fixed in the screen)
    for g in range(par_ng):
        for r in range(ps):
            offg[g] += Zp[g, r] * par_w[r + 1]
    S_buf = np.empty(ps + 1, np.int64)
    w_buf = np.zeros(ps + 2)
    Z_buf = np.empty((0, 1))
    cw_buf = np.empty(0)
    lo_b = np.full(ps + 2, -bound)
    hi_b = np.full(ps + 2, bound)
    lo_b[0] = -1e300
    hi_b[0] = 1e300
    a0s = np.empty(par_ng)
    a1s = np.empty(par_ng)
    th = np.empty(ps + 2)
    u = 0
    while u < m:
        v = var_of[js[u]]
        u2 = u
        while u2 < m and var_of[js[u2]] == v:
            u2 += 1
        nc = ncode[v]
        H = np.zeros((par_ng, nc))
        if kind[v] == 0 and (u2 - u) * n < n + par_ng * nc:
            # few thresholds of this chain: accumulate the weight of x_j = 1 per group directly
            for t_ in range(vrp[v], vrp[v + 1]):
                i = vrows[t_]
                qi = QT[v, i]
                g = par_inv[i]
                for t in range(u, u2):
                    l = lev_of[js[t]]
                    if qi >= l:
                        H[g, l] += c[i]
        else:
            for t_ in range(vrp[v], vrp[v + 1]):
                i = vrows[t_]
                H[par_inv[i], QT[v, i]] += c[i]
        if kind[v] == 0 and not ((u2 - u) * n < n + par_ng * nc):
            for g in range(par_ng):
                for q in range(nc - 2, -1, -1):
                    H[g, q] += H[g, q + 1]
        for t in range(u, u2):
            j = js[t]
            S = S_buf
            w = w_buf
            w[:] = 0.0
            w[0] = par_w[0]
            pos = 0
            q = 0
            ins = False
            for r in range(ps + 1):
                if not ins and (q >= ps or j < par_S[q]):
                    S[r] = j
                    pos = r
                    w[r + 1] = init[t]
                    ins = True
                else:
                    S[r] = par_S[q]
                    w[r + 1] = par_w[q + 1]
                    q += 1
            if kind[v] == 0 and not screen:
                l = lev_of[j]
                for g in range(par_ng):
                    w1 = H[g, l]
                    a1s[g] = w1
                    w0 = par_gc[g] - w1
                    a0s[g] = w0 if w0 > 0.5 else 0.0
                th[0] = w[0]
                for r in range(ps):
                    th[r + 1] = par_w[r + 1]
                th[ps + 1] = init[t]
                loss = newton_split(Zp, par_gy, a0s, a1s, th, bound, 50, tol)
                w[0] = th[0]
                for r in range(ps + 1):
                    if r < pos:
                        w[r + 1] = th[r + 1]
                    elif r == pos:
                        w[r + 1] = th[ps + 1]
                    else:
                        w[r + 1] = th[r]
                losses[t] = loss
                Wout[t] = w
                continue
            maxcell = par_ng * (2 if kind[v] == 0 else nc)
            ncol = 3 if screen else ps + 2
            if Z_buf.shape[0] < maxcell or Z_buf.shape[1] != ncol:
                Z_buf = np.empty((maxcell, ncol))
                cw_buf = np.empty(maxcell)
            Z = Z_buf
            cw = cw_buf
            nce = 0
            for g in range(par_ng):
                yg = par_gy[g]
                ncg = 2 if kind[v] == 0 else nc
                for xq in range(ncg):
                    if kind[v] == 0:
                        w1 = H[g, lev_of[j]]
                        if xq == 1:
                            wt = w1
                        else:
                            w0 = par_gc[g] - w1
                            wt = w0 if w0 > 0.5 else 0.0
                        xv = float(xq)
                    else:
                        wt = H[g, xq]
                        xv = tabv[tptr[j] + xq]
                    if wt <= 0.0:
                        continue
                    Z[nce, 0] = yg
                    if screen:
                        Z[nce, 1] = yg * xv
                        Z[nce, 2] = offg[g]
                    else:
                        for r in range(ps + 1):
                            if r == pos:
                                Z[nce, r + 1] = yg * xv
                            else:
                                Z[nce, r + 1] = Zp[g, r if r < pos else r - 1]
                    cw[nce] = wt
                    nce += 1
            if screen:
                # only the intercept and the new coefficient move (an upper bound on the child's loss)
                w2 = np.array([w[0], 0.0, 1.0])
                lo2 = np.array([-1e300, -bound, 1.0])
                hi2 = np.array([1e300, bound, 1.0])
                loss, _ = newton_fit(Z[:nce], cw[:nce], w2, lo2, hi2, 20, tol)
                w[0] = w2[0]
                w[pos + 1] = w2[1]
            else:
                loss, _ = newton_fit(Z[:nce], cw[:nce], w, lo_b, hi_b, 50, tol)
            losses[t] = loss
            Wout[t] = w
        u = u2
    return losses, Wout
def chol_solve_buf(A, bb, nf, Lm, x)

chol_solve on the leading nf x nf block of A, with caller buffers Lm, x (x returned in x[:nf]).

Expand source code
@_kernel()
def chol_solve_buf(A, bb, nf, Lm, x):
    """chol_solve on the leading nf x nf block of A, with caller buffers Lm, x (x returned in x[:nf])."""
    for i in range(nf):
        for j in range(i + 1):
            sm = A[i, j]
            for t in range(j):
                sm -= Lm[i, t] * Lm[j, t]
            if i == j:
                if sm <= 0.0:
                    sol = np.linalg.solve(np.ascontiguousarray(A[:nf, :nf]), bb[:nf].copy())
                    for u in range(nf):
                        x[u] = sol[u]
                    return
                Lm[i, i] = np.sqrt(sm)
            else:
                Lm[i, j] = sm / Lm[j, j]
    for i in range(nf):
        x[i] = bb[i]
    for i in range(nf):
        sm = x[i]
        for t in range(i):
            sm -= Lm[i, t] * x[t]
        x[i] = sm / Lm[i, i]
    for i in range(nf - 1, -1, -1):
        sm = x[i]
        for t in range(i + 1, nf):
            sm -= Lm[t, i] * x[t]
        x[i] = sm / Lm[i, i]
def colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pow2, vrp, vrows)

out[j, p] = sum_i x_ij^(1 + pow2[p]) R[i, p] for every column, from per-code histograms of each variable.

Expand source code
@_kernel()
def colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pow2, vrp, vrows):
    """out[j, p] = sum_i x_ij^(1 + pow2[p]) R[i, p] for every column, from per-code histograms of each variable."""
    V, n = QT.shape
    P = R.shape[1]
    d = lev_of.shape[0]
    out = np.zeros((d, P))
    for v in range(V):
        nc = ncode[v]
        Bq = np.zeros((nc, P))
        for t_ in range(vrp[v], vrp[v + 1]):
            i = vrows[t_]
            q = QT[v, i]
            for p in range(P):
                Bq[q, p] += R[i, p]
        if kind[v] == 0:
            for q in range(nc - 2, -1, -1):  # suffix sums: column of level l sums codes >= l
                for p in range(P):
                    Bq[q, p] += Bq[q + 1, p]
            for t in range(vcp[v], vcp[v + 1]):
                j = vcols[t]
                for p in range(P):
                    out[j, p] = Bq[lev_of[j], p]
        else:
            for t in range(vcp[v], vcp[v + 1]):
                j = vcols[t]
                base = tptr[j]
                for p in range(P):
                    s = 0.0
                    for q in range(nc):
                        x = tabv[base + q]
                        s += (x * x if pow2[p] else x) * Bq[q, p]
                    out[j, p] = s
    return out
def complement_rep(B, n, cand, d, R)

rep[j]: the smaller index of a complement pair (x_j + x_j' = 1 on every row) of binary columns, else j. Columns are matched by a hash of their bitsets and of the complemented bitsets, then checked exactly.

Expand source code
@_kernel()
def complement_rep(B, n, cand, d, R):
    """rep[j]: the smaller index of a complement pair (x_j + x_j' = 1 on every row) of binary columns, else j.
    Columns are matched by a hash of their bitsets and of the complemented bitsets, then checked exactly."""
    nw = B.shape[0]
    rep = np.arange(d)
    last = n - (nw - 1) * 64
    lastmask = np.uint64(0xFFFFFFFFFFFFFFFF) if last == 64 else (np.uint64(1) << np.uint64(last)) - np.uint64(1)
    full = np.uint64(0xFFFFFFFFFFFFFFFF)
    h = Dict.empty(key_type=nb.types.uint64, value_type=nb.types.int64)
    for t in range(cand.shape[0]):
        j = cand[t]
        hv = np.uint64(0)
        for w in range(nw):
            hv += B[w, j] * R[w]
        h[hv] = j
    for t in range(cand.shape[0]):
        j = cand[t]
        hc = np.uint64(0)
        for w in range(nw):
            m = lastmask if w == nw - 1 else full
            hc += ((~B[w, j]) & m) * R[w]
        if hc in h:
            j2 = h[hc]
            ok = True
            for w in range(nw):
                m = lastmask if w == nw - 1 else full
                if ((~B[w, j]) & m) != B[w, j2]:
                    ok = False
                    break
            if ok:
                rep[j] = min(j, j2)
    return rep
def eval_cells(q, deltas, cb, cx, cwp, cwn, nt, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b, tlo, TB)

out[:, q, r]: estimated calibrated loss and loss at fixed (a, b) after s += deltas[q, r] * x, from the touched cells (bin cb, value cx, weights cwp / cwn of y = +1 / -1). The estimate is the exact loss at the current (a, b) minus a trust-region Newton step in (a, b).

Expand source code
@_kernel()
def eval_cells(q, deltas, cb, cx, cwp, cwn, nt, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b, tlo, TB):
    """out[:, q, r]: estimated calibrated loss and loss at fixed (a, b) after s += deltas[q, r] * x, from the
    touched cells (bin cb, value cx, weights cwp / cwn of y = +1 / -1). The estimate is the exact loss at the
    current (a, b) minus a trust-region Newton step in (a, b)."""
    nd = deltas.shape[1]
    wk = np.zeros((9, nd))  # one allocation for the nine work rows
    stepa = wk[0]
    stepb = wk[1]
    hyb = wk[2]
    dL = wk[3]
    dGa = wk[4]
    dGb = wk[5]
    dHaa = wk[6]
    dHab = wk[7]
    dHbb = wk[8]
    oL = oGa = oGb = oHaa = oHab = oHbb = 0.0
    for u in range(nt):
        bq = cb[u]
        x = cx[u]
        wpc = cwp[u]
        wnc = cwn[u]
        so = sv[bq]
        lo = wpc * lp[bq] + wnc * ln_[bq]
        go = -wpc * pp[bq] + wnc * pn[bq]
        wo = (wpc + wnc) * wq[bq]
        oL += lo
        oGa += go * so
        oGb += go
        oHaa += wo * so * so
        oHab += wo * so
        oHbb += wo
        for r in range(nd):
            sn_ = so + deltas[q, r] * x
            tf = sn_ - tlo
            if TB.shape[1] > 0 and tf == np.floor(tf) and tf >= 0.0 and tf < TB.shape[1]:
                # integer scores: the loss terms at (a, b) come from a table over the score range
                ti = int(tf)
                lpn = TB[0, ti]
                lnn = TB[1, ti]
                ppn = TB[2, ti]
                pnn = TB[3, ti]
            else:
                z = a * sn_ + b
                e = np.exp(-abs(z))
                l1 = np.log1p(e)
                if z > 0:
                    lpn = l1
                    lnn = z + l1
                    ppn = e / (1.0 + e)
                    pnn = 1.0 / (1.0 + e)
                else:
                    lpn = -z + l1
                    lnn = l1
                    ppn = 1.0 / (1.0 + e)
                    pnn = e / (1.0 + e)
            gn = -wpc * ppn + wnc * pnn
            wn = (wpc + wnc) * ppn * pnn
            dL[r] += wpc * lpn + wnc * lnn - lo
            dGa[r] += gn * sn_ - go * so
            dGb[r] += gn - go
            dHaa[r] += wn * sn_ * sn_ - wo * so * so
            dHab[r] += wn * sn_ - wo * so
            dHbb[r] += wn - wo
    for r in range(nd):
        ga = T[1] + dGa[r]
        gb = T[2] + dGb[r]
        haa = T[3] + dHaa[r] + 1e-12
        hab = T[4] + dHab[r]
        hbb = T[5] + dHbb[r] + 1e-12
        det = haa * hbb - hab * hab
        if det > 1e-14 * haa * hbb:
            da = (hbb * ga - hab * gb) / det
            db = (haa * gb - hab * ga) / det
        else:
            da = 0.0
            db = gb / hbb
        t = 1.0
        if abs(da) * t > tr_a * abs(a) + 1e-300:
            t = tr_a * abs(a) / abs(da)
        if abs(db) * t > tr_b:
            t = tr_b / abs(db)
        stepa[r] = -t * da
        stepb[r] = -t * db
        out[1, q, r] = T[0] + dL[r]
    # the untouched rows at the Newton point, by their quadratic model around (a, b)
    UGa = T[1] - oGa
    UGb = T[2] - oGb
    UHaa = T[3] - oHaa
    UHab = T[4] - oHab
    UHbb = T[5] - oHbb
    for r in range(nd):
        hyb[r] = (T[0] - oL + UGa * stepa[r] + UGb * stepb[r]
                  + 0.5 * (UHaa * stepa[r] * stepa[r] + 2 * UHab * stepa[r] * stepb[r] + UHbb * stepb[r] * stepb[r]))
    # the touched cells exactly at the Newton point
    for u in range(nt):
        so = sv[cb[u]]
        x = cx[u]
        wpc = cwp[u]
        wnc = cwn[u]
        for r in range(nd):
            z = (a + stepa[r]) * (so + deltas[q, r] * x) + b + stepb[r]
            l1 = np.log1p(np.exp(-abs(z)))
            if z > 0:
                hyb[r] += wpc * l1 + wnc * (z + l1)
            else:
                hyb[r] += wpc * (l1 - z) + wnc * l1
    for r in range(nd):
        out[0, q, r] = min(out[1, q, r], hyb[r])
def eval_support(S, deltas, HP, HN, lev_of, nb_, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b)

Value changes of the support columns (all chain columns) from the support histograms.

Expand source code
@_kernel()
def eval_support(S, deltas, HP, HN, lev_of, nb_, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b):
    """Value changes of the support columns (all chain columns) from the support histograms."""
    m = S.shape[0]
    dlo = 0.0
    dhi = 0.0
    integral = True
    for q in range(m):
        for r in range(deltas.shape[1]):
            dv = deltas[q, r]
            if dv != np.floor(dv):
                integral = False
            dlo = min(dlo, dv)
            dhi = max(dhi, dv)
    if integral:
        tlo, TB = make_tb(sv, nb_, dlo, dhi, a, b)
    else:
        tlo, TB = 0.0, np.zeros((4, 0))
    cb = np.empty(nb_, np.int64)
    cx = np.ones(nb_)
    cwp = np.empty(nb_)
    cwn = np.empty(nb_)
    for q in range(m):
        l = lev_of[S[q]]
        nt = 0
        for bq in range(nb_):
            if HP[q, bq, l] > 0.0 or HN[q, bq, l] > 0.0:
                cb[nt] = bq
                cwp[nt] = HP[q, bq, l]
                cwn[nt] = HN[q, bq, l]
                nt += 1
        eval_cells(q, deltas, cb, cx, cwp, cwn, nt, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b, tlo, TB)
def eval_vars(cols, qidx, deltas, inv, nbins, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind, ncode, tptr, tabv, out, tr_a, tr_b, vrp, vrows)

Moves s += deltas[q] * x_cols[q] for columns sorted by variable (qidx: their rows in deltas / out). Rows are histogrammed once per variable by (score bin, code); a chain's columns read suffix sums.

Expand source code
@_kernel()
def eval_vars(cols, qidx, deltas, inv, nbins, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind,
              ncode, tptr, tabv, out, tr_a, tr_b, vrp, vrows):
    """Moves s += deltas[q] * x_cols[q] for columns sorted by variable (qidx: their rows in deltas / out).
    Rows are histogrammed once per variable by (score bin, code); a chain's columns read suffix sums."""
    m = cols.shape[0]
    n = inv.shape[0]
    # table of the per-score loss terms at (a, b) when scores and moves are integers
    tlo = 0.0
    TB = np.zeros((4, 0))
    integral = True
    lo = np.inf
    hi = -np.inf
    for q in range(nbins):
        if sv[q] != np.floor(sv[q]):
            integral = False
        lo = min(lo, sv[q])
        hi = max(hi, sv[q])
    dlo = 0.0
    dhi = 0.0
    for t in range(m):
        for r in range(deltas.shape[1]):
            dv = deltas[qidx[t], r]
            if dv != np.floor(dv):
                integral = False
            dlo = min(dlo, dv)
            dhi = max(dhi, dv)
    if integral and nbins > 0 and hi - lo + dhi - dlo < 4096:
        tlo = lo + dlo
        R = int(hi + dhi - tlo) + 1
        TB = np.empty((4, R))
        for ti in range(R):
            z = a * (tlo + ti) + b
            e = np.exp(-abs(z))
            l1 = np.log1p(e)
            if z > 0:
                TB[0, ti] = l1
                TB[1, ti] = z + l1
                TB[2, ti] = e / (1.0 + e)
                TB[3, ti] = 1.0 / (1.0 + e)
            else:
                TB[0, ti] = -z + l1
                TB[1, ti] = l1
                TB[2, ti] = 1.0 / (1.0 + e)
                TB[3, ti] = e / (1.0 + e)
    u = 0
    while u < m:
        v = var_of[cols[u]]
        u2 = u
        while u2 < m and var_of[cols[u2]] == v:
            u2 += 1
        nc = ncode[v]
        kv = kind[v]
        maxc = nbins if kv == 0 else (nbins * nc if kv == 1 else n)
        cb = np.empty(maxc, np.int64)
        cx = np.empty(maxc)
        cwp = np.empty(maxc)
        cwn = np.empty(maxc)
        if kv < 2:
            HP = np.zeros((nbins, nc))
            HN = np.zeros((nbins, nc))
            for t_ in range(vrp[v], vrp[v + 1]):
                i = vrows[t_]
                if y[i] > 0:
                    HP[inv[i], QT[v, i]] += c[i]
                else:
                    HN[inv[i], QT[v, i]] += c[i]
            if kv == 0:
                for bq in range(nbins):
                    for qq in range(nc - 2, -1, -1):
                        HP[bq, qq] += HP[bq, qq + 1]
                        HN[bq, qq] += HN[bq, qq + 1]
        for t in range(u, u2):
            j = cols[t]
            nt = 0
            if kv == 0:
                l = lev_of[j]
                for bq in range(nbins):
                    if HP[bq, l] > 0.0 or HN[bq, l] > 0.0:
                        cb[nt] = bq
                        cx[nt] = 1.0
                        cwp[nt] = HP[bq, l]
                        cwn[nt] = HN[bq, l]
                        nt += 1
            elif kv == 1:
                for bq in range(nbins):
                    for qq in range(nc):
                        x = tabv[tptr[j] + qq]
                        if x != 0.0 and (HP[bq, qq] > 0.0 or HN[bq, qq] > 0.0):
                            cb[nt] = bq
                            cx[nt] = x
                            cwp[nt] = HP[bq, qq]
                            cwn[nt] = HN[bq, qq]
                            nt += 1
            else:
                for i in range(n):
                    x = tabv[tptr[j] + QT[v, i]]
                    if x != 0.0:
                        cb[nt] = inv[i]
                        cx[nt] = x
                        cwp[nt] = c[i] if y[i] > 0 else 0.0
                        cwn[nt] = c[i] - cwp[nt]
                        nt += 1
            eval_cells(qidx[t], deltas, cb, cx, cwp, cwn, nt, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b, tlo, TB)
        u = u2
    return out
def exact_moves(s, w, S, free, k, y, c, a, b, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp, nzv, n_screen, L0)

Exact best swap / addition: per removal (and no removal when |S| < k) the n_screen best columns by the continuous Newton gain at the removal state with a refitted (a, b) (one colsum for all removals); the POLV best values of every screened chain column (by the loss at the refitted (a, b)) scored by their exact calibrated loss on (bin, x) cells. Returns (total loss, removed column or -1, added column, value); added column -1 when nothing beats L0.

Expand source code
@_kernel()
def exact_moves(s, w, S, free, k, y, c, a, b, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp,
                nzv, n_screen, L0):
    """Exact best swap / addition: per removal (and no removal when |S| < k) the n_screen best columns by the
    continuous Newton gain at the removal state with a refitted (a, b) (one colsum for all removals); the POLV best
    values of every screened chain column (by the loss at the refitted (a, b)) scored
    by their exact calibrated loss on (bin, x) cells. Returns (total loss, removed column or -1, added column,
    value); added column -1 when nothing beats L0."""
    n = s.shape[0]
    m = S.shape[0]
    d = w.shape[0]
    nv = nzv.shape[0]
    vb = 0.0
    for r in range(nv):
        vb = max(vb, abs(nzv[r]))
    bL = L0
    brj = -1
    baj = -1
    bv = 0.0
    nfree = 0
    for t in range(d):
        if free[t] and kind[var_of[t]] == 0:
            nfree += 1
    if nfree == 0:
        return bL, brj, baj, bv
    fcols = np.empty(nfree, np.int64)
    u = 0
    for t in range(d):
        if free[t] and kind[var_of[t]] == 0:
            fcols[u] = t
            u += 1
    q0 = -1 if m < k else 0
    nq = m - q0
    R = np.empty((n, 2 * nq))
    INV = np.empty((nq, n), np.int64)
    BINS = List()
    AB = np.empty((nq, 2))
    s2 = np.empty(n)
    for qi in range(nq):
        q = q0 + qi
        if q >= 0:
            j = S[q]
            base = tptr[j]
            vj = var_of[j]
            for i in range(n):
                s2[i] = s[i] - w[j] * tabv[base + QT[vj, i]]
        else:
            for i in range(n):
                s2[i] = s[i]
        inv, sv, Wp, Wn = bin_scores(s2, y, c)
        nb_ = sv.shape[0]
        if nb_ >= 2:
            _, aq, bq = calibrate_bins(sv, Wp, Wn, a, b)
        else:
            aq, bq = a, b
        AB[qi, 0] = aq
        AB[qi, 1] = bq
        INV[qi] = inv
        st3 = np.empty((3, nb_))
        st3[0] = sv
        st3[1] = Wp
        st3[2] = Wn
        BINS.append(st3)
        lp = np.empty(nb_)
        ln_ = np.empty(nb_)
        pp = np.empty(nb_)
        pn = np.empty(nb_)
        wq = np.empty(nb_)
        bin_stats(sv, Wp, Wn, aq, bq, lp, ln_, pp, pn, wq)
        for i in range(n):
            g = inv[i]
            R[i, 2 * qi] = (-pp[g] if y[i] > 0 else pn[g]) * c[i]
            R[i, 2 * qi + 1] = wq[g] * c[i]
    pw = np.zeros(2 * nq, np.bool_)
    for qi in range(nq):
        pw[2 * qi + 1] = True
    G = colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pw, vrp, vrows)
    fl = np.zeros(nv)
    for qi in range(nq):
        q = q0 + qi
        j = S[q] if q >= 0 else -1
        aq = AB[qi, 0]
        bq = AB[qi, 1]
        st3 = BINS[qi]
        sv = st3[0]
        Wp = st3[1]
        Wn = st3[2]
        nb_ = sv.shape[0]
        inv = INV[qi]
        cols = screen_newton(G, 2 * qi, fcols, n_screen)
        cwp0 = np.empty(nb_)
        cwn0 = np.empty(nb_)
        cwp1 = np.empty(nb_)
        cwn1 = np.empty(nb_)
        svq = np.empty(2 * nb_)
        Wpq = np.empty(2 * nb_)
        Wnq = np.empty(2 * nb_)
        E = np.empty(4 * nb_)
        En = np.empty(4 * nb_)
        tlo, TB = make_tb(sv, nb_, -vb, vb, aq, bq)
        # columns grouped by variable: one (bin, code) histogram per variable, suffix-summed over codes
        vo = np.empty(cols.shape[0], np.int64)
        for t in range(cols.shape[0]):
            vo[t] = var_of[cols[t]] * (d + 1) + lev_of[cols[t]]
        cols = cols[np.argsort(vo)]
        cur_v = -1
        HPv = np.zeros((1, 1))
        HNv = np.zeros((1, 1))
        for t in range(cols.shape[0]):
            l = cols[t]
            if l == j:
                continue
            v = var_of[l]
            lv = lev_of[l]
            if v != cur_v:
                cur_v = v
                nc = ncode[v]
                HPv = np.zeros((nc + 1, nb_))
                HNv = np.zeros((nc + 1, nb_))
                for tt in range(vrp[v], vrp[v + 1]):
                    i = vrows[tt]
                    if y[i] > 0:
                        HPv[QT[v, i], inv[i]] += c[i]
                    else:
                        HNv[QT[v, i], inv[i]] += c[i]
                for cc in range(nc - 1, 0, -1):
                    for g in range(nb_):
                        HPv[cc, g] += HPv[cc + 1, g]
                        HNv[cc, g] += HNv[cc + 1, g]
            for g in range(nb_):
                cwp1[g] = HPv[lv, g]
                cwn1[g] = HNv[lv, g]
            for g in range(nb_):
                cwp0[g] = Wp[g] - cwp1[g]
                cwn0[g] = Wn[g] - cwn1[g]
            fl[:] = 0.0
            for g in range(nb_):
                if cwp1[g] <= 0.0 and cwn1[g] <= 0.0:
                    continue
                for r in range(nv):
                    tf = sv[g] + nzv[r] - tlo
                    if TB.shape[1] > 0 and tf >= 0.0 and tf < TB.shape[1] and tf == np.floor(tf):
                        ti = int(tf)
                        fl[r] += cwp1[g] * TB[0, ti] + cwn1[g] * TB[1, ti]
                        continue
                    z = aq * (sv[g] + nzv[r]) + bq
                    l1 = np.log1p(np.exp(-abs(z)))
                    if z > 0:
                        fl[r] += cwp1[g] * l1 + cwn1[g] * (z + l1)
                    else:
                        fl[r] += cwp1[g] * (l1 - z) + cwn1[g] * l1
            fo = np.argsort(fl)
            for r0 in range(min(POLV, nv)):
                r = fo[r0]
                L2 = _merge_calib(sv, cwp0, cwn0, cwp1, cwn1, nzv[r], aq, bq, svq, Wpq, Wnq, E, En, bL)
                if L2 < bL:
                    bL, brj, baj, bv = L2, j, l, nzv[r]
    return bL, brj, baj, bv
def exact_values(w, S, s, L, a, b, inv, sv, Wp, Wn, y, c, QT, var_of, kind, ncode, lev_of, tptr, tabv, allv)

Exact calibrated loss of every value change of every support column (chain columns from one (bin, code) histogram pass; other columns by a row pass each). Returns (best total loss, column, value); column -1 if none improves.

Expand source code
@_kernel()
def exact_values(w, S, s, L, a, b, inv, sv, Wp, Wn, y, c, QT, var_of, kind, ncode, lev_of, tptr, tabv, allv):
    """Exact calibrated loss of every value change of every support column (chain columns from one (bin, code)
    histogram pass; other columns by a row pass each). Returns (best total loss, column, value); column -1 if none
    improves."""
    nb_ = sv.shape[0]
    SHP, SHN = support_hist(S, inv, nb_, y, c, QT, var_of, kind, ncode, lev_of, 0)
    bL = L - 1e-9 * L
    bj = -1
    bv = 0.0
    m2 = 2 * nb_
    cwp = np.empty(m2)
    cwn = np.empty(m2)
    svq = np.empty(m2)
    Wpq = np.empty(m2)
    Wnq = np.empty(m2)
    E = np.empty(2 * m2)
    En = np.empty(2 * m2)
    n = s.shape[0]
    for q in range(S.shape[0]):
        j = S[q]
        v = var_of[j]
        if kind[v] == 0:
            l = lev_of[j]
            for bq in range(nb_):
                cwp[2 * bq] = Wp[bq] - SHP[q, bq, l]
                cwn[2 * bq] = Wn[bq] - SHN[q, bq, l]
                cwp[2 * bq + 1] = SHP[q, bq, l]
                cwn[2 * bq + 1] = SHN[q, bq, l]
            c0p = np.empty(nb_)
            c0n = np.empty(nb_)
            c1p = np.empty(nb_)
            c1n = np.empty(nb_)
            for bq in range(nb_):
                c0p[bq] = cwp[2 * bq]
                c0n[bq] = cwn[2 * bq]
                c1p[bq] = cwp[2 * bq + 1]
                c1n[bq] = cwn[2 * bq + 1]
            # loss at the current (a, b) of every value; the POLVV best calibrated exactly
            fl = np.full(allv.shape[0], np.inf)
            for r in range(allv.shape[0]):
                dv = allv[r] - w[j]
                if dv == 0.0:
                    continue
                t = 0.0
                for bq in range(nb_):
                    if c1p[bq] <= 0.0 and c1n[bq] <= 0.0:
                        continue
                    z = a * (sv[bq] + dv) + b
                    l1 = np.log1p(np.exp(-abs(z)))
                    if z > 0:
                        t += c1p[bq] * l1 + c1n[bq] * (z + l1)
                    else:
                        t += c1p[bq] * (l1 - z) + c1n[bq] * l1
                fl[r] = t
            fo = np.argsort(fl)
            for r0 in range(min(POLVV, allv.shape[0] - 1)):
                r = fo[r0]
                dv = allv[r] - w[j]
                L2 = _merge_calib(sv, c0p, c0n, c1p, c1n, dv, a, b, svq, Wpq, Wnq, E, En, bL)
                if L2 < bL:
                    bL, bj, bv = L2, j, allv[r]
        else:
            base = tptr[j]
            s2 = np.empty(n)
            for r in range(allv.shape[0]):
                dv = allv[r] - w[j]
                if dv == 0.0:
                    continue
                for i in range(n):
                    s2[i] = s[i] + dv * tabv[base + QT[v, i]]
                inv2, sv2, Wp2, Wn2 = bin_scores(s2, y, c)
                if sv2.shape[0] < 2:
                    continue
                L2, a2, b2 = calibrate_bins(sv2, Wp2, Wn2, a, b)
                if L2 < bL:
                    bL, bj, bv = L2, j, allv[r]
    return bL, bj, bv
def gather_cols_nz(Q, idx, kind)

Q[:, idx] and, per variable, the CSR list of rows that can have a nonzero value (code > 0 for a chain, every row otherwise), in two passes (counts while gathering).

Expand source code
@_kernel()
def gather_cols_nz(Q, idx, kind):
    """Q[:, idx] and, per variable, the CSR list of rows that can have a nonzero value (code > 0 for a chain,
    every row otherwise), in two passes (counts while gathering)."""
    V = Q.shape[0]
    m = idx.shape[0]
    out = np.empty((V, m), Q.dtype)
    vrp = np.zeros(V + 1, np.int64)
    for v in range(V):
        cnt = 0
        for t in range(m):
            q = Q[v, idx[t]]
            out[v, t] = q
            cnt += q > 0
        vrp[v + 1] = vrp[v] + (m if kind[v] > 0 else cnt)
    vrows = np.empty(vrp[V], np.int32)
    for v in range(V):
        t0 = vrp[v]
        if kind[v] > 0:
            for i in range(m):
                vrows[t0 + i] = i
        else:
            for i in range(m):
                if out[v, i] > 0:
                    vrows[t0] = i
                    t0 += 1
    return out, vrp, vrows
def group_design(QT, var_of, tptr, tabv, gy, rep, S)
Expand source code
@_kernel()
def group_design(QT, var_of, tptr, tabv, gy, rep, S):
    ng = rep.shape[0]
    p = S.shape[0] + 1
    Z = np.empty((ng, p))
    for g in range(ng):
        Z[g, 0] = gy[g]
        for q in range(1, p):
            Z[g, q] = gy[g] * xval(QT, var_of, tptr, tabv, rep[g], S[q - 1])
    return Z
def ils_kernel(w, s, L, a, b, inv, sv, Wp, Wn, visited, hw, k, max_iter, swaps, gate, N, ne, n_screen, add_screen, margin, y, c, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp, valid, allv, nzv, tr_a, tr_b)

Best-improvement local search on the integer lattice (value changes and additions; swaps when those fail), candidates ranked by the calibrated-loss estimate and the best 2 ne checked exactly. visited: keys of points already expanded in this fit (the search from them is deterministic). Returns (loss, w).

Expand source code
@_kernel()
def ils_kernel(w, s, L, a, b, inv, sv, Wp, Wn, visited, hw, k, max_iter, swaps, gate, N, ne, n_screen, add_screen,
               margin, y, c, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp, valid,
               allv, nzv, tr_a, tr_b):
    """Best-improvement local search on the integer lattice (value changes and additions; swaps when those fail),
    candidates ranked by the calibrated-loss estimate and the best 2 ne checked exactly. visited: keys of points
    already expanded in this fit (the search from them is deterministic). Returns (loss, w)."""
    d = w.shape[0]
    m_chk = 2 * ne
    # the candidates of a failing swap phase at the final point (reused by the refit swaps)
    has_c = False
    c_oe = np.zeros(0)
    c_orj = np.zeros(0, np.int64)
    c_oaj = np.zeros(0, np.int64)
    c_ov = np.zeros(0)
    for it in range(max_iter):
        key = 0.0
        for j in range(d):
            if w[j] != 0.0:
                key += w[j] * hw[j]
        if swaps:
            if key in visited:
                break
            visited[key] = True
        nb_ = sv.shape[0]
        lp = np.empty(nb_)
        ln_ = np.empty(nb_)
        pp = np.empty(nb_)
        pn = np.empty(nb_)
        wq = np.empty(nb_)
        T = bin_stats(sv, Wp, Wn, a, b, lp, ln_, pp, pn, wq)
        cntS = 0
        for j in range(d):
            if w[j] != 0.0:
                cntS += 1
        S = np.empty(cntS, np.int64)
        free = np.empty(d, np.bool_)
        u = 0
        nfree = 0
        for j in range(d):
            if w[j] != 0.0:
                S[u] = j
                u += 1
            free[j] = w[j] == 0.0 and valid[j]
            nfree += free[j]
        SHP, SHN = support_hist(S, inv, nb_, y, c, QT, var_of, kind, ncode, lev_of, SLIDER)
        oe, oaj, ov, ofx = main_phase(w, S, free, k, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of,
                                      lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp, allv, nzv,
                                      ne, tr_a, tr_b, add_screen, SHP, SHN, Wp, Wn)
        orj = np.full(oe.shape[0], -1, np.int64)
        sel = top_order(oe, m_chk)
        bi, L2, s2, inv2, sv2, Wp2, Wn2, a2, b2 = check_kernel(s, w, L, oe[sel], orj[sel], oaj[sel], ov[sel], a, b,
                                                               margin, y, c, QT, var_of, tptr, tabv, kind, lev_of,
                                                               inv, sv)
        rj = -1
        aj = -1
        v = 0.0
        from_main = False
        if bi >= 0:
            rj = orj[sel[bi]]
            aj = oaj[sel[bi]]
            v = ov[sel[bi]]
            from_main = True
            m_sel = sel
            m_bi = bi
        if bi < 0 and cntS > 0 and nfree > 0 and swaps and L / N <= gate:
            # threshold slides before the full swap phase
            se, srj, saj, sv_ = slide_moves(w, S, free, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of,
                                            lev_of, kind, ncode, vrp, vrows, vcols, vcp, ne, tr_a, tr_b,
                                            SLIDER, SHP, SHN)
            sel = top_order(se, m_chk)
            bi, L2, s2, inv2, sv2, Wp2, Wn2, a2, b2 = check_kernel(s, w, L, se[sel], srj[sel], saj[sel], sv_[sel],
                                                                   a, b, margin, y, c, QT, var_of, tptr, tabv, kind,
                                                                   lev_of, inv, sv)
            if bi >= 0:
                rj = srj[sel[bi]]
                aj = saj[sel[bi]]
                v = sv_[sel[bi]]
        if bi >= 0:
            pass
        elif nfree > 0 and swaps and L / N <= gate:
            oe, orj, oaj, ov, ofx = swap_phase(s, w, S, free, y, c, a, b, QT, var_of, lev_of, kind, ncode, tptr, tabv,
                                               vrp, vrows, vcols, vcp, nzv, n_screen, ne, tr_a, tr_b, inv, sv)
            sel = top_order(oe, m_chk)
            bi, L2, s2, inv2, sv2, Wp2, Wn2, a2, b2 = check_kernel(s, w, L, oe[sel], orj[sel], oaj[sel], ov[sel], a,
                                                                   b, margin, y, c, QT, var_of, tptr, tabv, kind,
                                                                   lev_of, inv, sv)
            if bi >= 0:
                rj = orj[sel[bi]]
                aj = oaj[sel[bi]]
                v = ov[sel[bi]]
            else:
                has_c = True
                c_oe, c_orj, c_oaj, c_ov = oe, orj, oaj, ov
        if bi < 0:
            break
        if rj >= 0:
            w[rj] = 0.0
        w[aj] = v
        s, L, a, b, inv, sv, Wp, Wn = s2, L2, a2, b2, inv2, sv2, Wp2, Wn2
        if from_main:
            # the other checked value changes (other columns), exactly on the new point, before a new main phase
            used = np.zeros(1, np.int64)
            used[0] = aj
            for t in range(m_sel.shape[0]):
                if t == m_bi:
                    continue
                q = m_sel[t]
                cj = oaj[q]
                if orj[q] >= 0 or w[cj] == 0.0 or cj == aj:
                    continue  # value changes of support columns only
                dup = False
                for u in range(used.shape[0]):
                    if used[u] == cj:
                        dup = True
                if dup:
                    continue
                one_e = np.full(1, -np.inf)
                bi3, L3, s3, inv3, sv3, Wp3, Wn3, a3, b3 = check_kernel(s, w, L, one_e, orj[q:q + 1], oaj[q:q + 1],
                                                                       ov[q:q + 1], a, b, margin, y, c, QT, var_of,
                                                                       tptr, tabv, kind, lev_of, inv, sv)
                if bi3 >= 0:
                    w[cj] = ov[q]
                    s, L, a, b, inv, sv, Wp, Wn = s3, L3, a3, b3, inv3, sv3, Wp3, Wn3
                    used = np.append(used, cj)
    return L / N, w, has_c, c_oe, c_orj, c_oaj, c_ov
def main_phase(w, S, free, k, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp, allv, nzv, ne, tr_a, tr_b, add_screen, SHP, SHN, Wp, Wn)

Value changes of the support columns (every other value; 0 removes) and additions (when |S| < k): the ne best by the estimate per support column, and the ne best additions overall.

Expand source code
@_kernel()
def main_phase(w, S, free, k, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind, ncode, tptr,
               tabv, vrp, vrows, vcols, vcp, allv, nzv, ne, tr_a, tr_b, add_screen, SHP, SHN, Wp, Wn):
    """Value changes of the support columns (every other value; 0 removes) and additions (when |S| < k): the ne
    best by the estimate per support column, and the ne best additions overall."""
    m = S.shape[0]
    d = w.shape[0]
    nv = nzv.shape[0]
    nb_ = sv.shape[0]
    o_est = np.full(m * ne + ne, np.inf)
    o_fx = np.full(m * ne + ne, np.inf)
    o_aj = np.full(m * ne + ne, -1, np.int64)
    o_v = np.zeros(m * ne + ne)
    if m > 0:
        newv = np.empty((m, allv.shape[0] - 1))
        deltas = np.empty((m, allv.shape[0] - 1))
        vo = np.empty(m, np.int64)
        for q in range(m):
            u = 0
            wj = w[S[q]]
            for r in range(allv.shape[0]):
                if allv[r] != wj:
                    newv[q, u] = allv[r]
                    deltas[q, u] = allv[r] - wj
                    u += 1
            vo[q] = var_of[S[q]]
        o = np.argsort(vo, kind="mergesort")
        out = np.empty((2, m, newv.shape[1]))
        allch = True
        for q in range(m):
            if kind[var_of[S[q]]] != 0:
                allch = False
        if allch:
            eval_support(S, deltas, SHP, SHN, lev_of, nb_, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b)
        else:
            eval_vars(S[o], o, deltas, inv, nb_, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind,
                      ncode, tptr, tabv, out, tr_a, tr_b, vrp, vrows)
        for q in range(m):
            fo = np.argsort(out[0, q], kind="mergesort")
            u = 0
            for t in range(fo.shape[0]):
                if u >= ne:
                    break
                r = fo[t]
                if not np.isfinite(out[0, q, r]):
                    continue
                o_est[q * ne + u] = out[0, q, r]
                o_fx[q * ne + u] = out[1, q, r]
                o_aj[q * ne + u] = S[q]
                o_v[q * ne + u] = newv[q, r]
                u += 1
    if m < k:
        cnt = 0
        for t in range(d):
            if free[t]:
                cnt += 1
        if cnt > 0:
            cols = np.empty(cnt, np.int64)
            vo = np.empty(cnt, np.int64)
            cnt = 0
            for t in range(d):
                if free[t]:
                    cols[cnt] = t
                    vo[cnt] = var_of[t]
                    cnt += 1
            if add_screen > 0 and cnt > add_screen:
                # second-order screen of the additions at the current (a, b); only the best are estimated
                R2 = np.empty((y.shape[0], 2))
                for i in range(y.shape[0]):
                    q = inv[i]
                    R2[i, 0] = (-pp[q] if y[i] > 0 else pn[q]) * c[i]
                    R2[i, 1] = wq[q] * c[i]
                pw2 = np.zeros(2, np.bool_)
                pw2[1] = True
                G12 = colsum(QT, R2, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pw2, vrp, vrows)
                cols = screen_cols(G12, 0, cols, a, nzv, add_screen)
                cnt = cols.shape[0]
                vo = np.empty(cnt, np.int64)
                for t in range(cnt):
                    vo[t] = var_of[cols[t]]
            o = np.argsort(vo, kind="mergesort")
            deltas = np.empty((cnt, nv))
            for t in range(cnt):
                deltas[t, :] = nzv
            out = np.empty((2, cnt, nv))
            eval_vars(cols[o], o, deltas, inv, nb_, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of,
                      kind, ncode, tptr, tabv, out, tr_a, tr_b, vrp, vrows)
            flat = out[0].ravel()
            fl1 = out[1].ravel()
            # the ne smallest finite estimates
            u = 0
            for _ in range(ne):
                bi = -1
                bv = np.inf
                for f in range(flat.shape[0]):
                    if flat[f] < bv:
                        dup = False
                        for z in range(u):
                            if o_aj[m * ne + z] == cols[f // nv] and o_v[m * ne + z] == nzv[f % nv]:
                                dup = True
                        if not dup:
                            bv = flat[f]
                            bi = f
                if bi < 0:
                    break
                o_est[m * ne + u] = flat[bi]
                o_fx[m * ne + u] = fl1[bi]
                o_aj[m * ne + u] = cols[bi // nv]
                o_v[m * ne + u] = nzv[bi % nv]
                u += 1
    return o_est, o_aj, o_v, o_fx
def make_child(par_inv, par_ng, par_S, par_w, j, colcode_j, ncode_j, y, c, QT, var_of, tptr, tabv, bound, tol)

Add column j to a parent: refine its row groups and refit (b0, beta_S) by projected Newton.

Expand source code
@_kernel()
def make_child(par_inv, par_ng, par_S, par_w, j, colcode_j, ncode_j, y, c, QT, var_of, tptr, tabv, bound, tol):
    """Add column j to a parent: refine its row groups and refit (b0, beta_S) by projected Newton."""
    inv, ng, gy, gc, rep = regroup(par_inv, par_ng, colcode_j, ncode_j, y, c)
    ps = par_S.shape[0]
    S = np.empty(ps + 1, np.int64)
    w = np.zeros(ps + 2)
    w[0] = par_w[0]
    q = 0
    ins = False
    for r in range(ps + 1):
        if not ins and (q >= ps or j < par_S[q]):
            S[r] = j
            w[r + 1] = 0.0
            ins = True
        else:
            S[r] = par_S[q]
            w[r + 1] = par_w[q + 1]
            q += 1
    lo = np.full(ps + 2, -bound)
    hi = np.full(ps + 2, bound)
    lo[0] = -1e300
    hi[0] = 1e300
    Z = group_design(QT, var_of, tptr, tabv, gy, rep, S)
    loss, mg = newton_fit(Z, gc, w, lo, hi, 50, tol)
    return S, inv, ng, gy, gc, rep, loss, mg, w
def make_tb(sv, nbins, dlo, dhi, a, b)

Loss terms at (a, b) tabulated over the integer scores sv + [dlo, dhi] (empty when scores are not integers).

Expand source code
@_kernel()
def make_tb(sv, nbins, dlo, dhi, a, b):
    """Loss terms at (a, b) tabulated over the integer scores sv + [dlo, dhi] (empty when scores are not integers)."""
    tlo = 0.0
    TB = np.zeros((4, 0))
    integral = True
    lo = np.inf
    hi = -np.inf
    for q in range(nbins):
        if sv[q] != np.floor(sv[q]):
            integral = False
        lo = min(lo, sv[q])
        hi = max(hi, sv[q])
    if dlo != np.floor(dlo) or dhi != np.floor(dhi):
        integral = False
    if integral and nbins > 0 and hi - lo + dhi - dlo < 4096:
        tlo = lo + dlo
        R = int(hi + dhi - tlo) + 1
        TB = np.empty((4, R))
        for ti in range(R):
            z = a * (tlo + ti) + b
            e = np.exp(-abs(z))
            l1 = np.log1p(e)
            if z > 0:
                TB[0, ti] = l1
                TB[1, ti] = z + l1
                TB[2, ti] = e / (1.0 + e)
                TB[3, ti] = 1.0 / (1.0 + e)
            else:
                TB[0, ti] = -z + l1
                TB[1, ti] = l1
                TB[2, ti] = 1.0 / (1.0 + e)
                TB[3, ti] = e / (1.0 + e)
    return tlo, TB
def materialise_kernel(par_inv, par_ng, j, S, w, QT, var_of, kind, lev_of, ncode, tptr, tabv, y, c)
Expand source code
@_kernel()
def materialise_kernel(par_inv, par_ng, j, S, w, QT, var_of, kind, lev_of, ncode, tptr, tabv, y, c):
    v = var_of[j]
    n = par_inv.shape[0]
    code = np.empty(n, np.int64)
    if kind[v] == 0:
        # binary column: the parent's groups split by x_j in one pass (groups numbered by first appearance)
        l = lev_of[j]
        table = np.full(2 * par_ng, -1, np.int64)
        inv = np.empty(n, np.int64)
        gy = np.empty(n)
        gc = np.zeros(n)
        rep = np.empty(n, np.int64)
        ng = 0
        for i in range(n):
            key = 2 * par_inv[i] + (1 if QT[v, i] >= l else 0)
            g = table[key]
            if g < 0:
                g = ng
                table[key] = g
                rep[g] = i
                gy[g] = y[i]
                ng += 1
            inv[i] = g
            gc[g] += c[i]
        gy = gy[:ng].copy()
        gc = gc[:ng].copy()
        rep = rep[:ng].copy()
        Z = group_design(QT, var_of, tptr, tabv, gy, rep, S)
        mg = np.zeros(ng)
        for g in range(ng):
            for q in range(Z.shape[1]):
                mg[g] += Z[g, q] * w[q]
        return inv, ng, gy, gc, rep, mg
    else:
        for i in range(n):
            code[i] = QT[v, i]
        nc = ncode[v]
    inv, ng, gy, gc, rep = regroup(par_inv, par_ng, code, nc, y, c)
    Z = group_design(QT, var_of, tptr, tabv, gy, rep, S)
    mg = np.zeros(ng)
    for g in range(ng):
        for q in range(Z.shape[1]):
            mg[g] += Z[g, q] * w[q]
    return inv, ng, gy, gc, rep, mg
def newton_fit(Z, c, w, lo, hi, maxit, tol)

Projected Newton for min sum_g c_g log(1 + exp(-Z_g . w)) with box bounds; w updated in place. Returns (loss, margins).

Expand source code
@_kernel()
def newton_fit(Z, c, w, lo, hi, maxit, tol):
    """Projected Newton for min sum_g c_g log(1 + exp(-Z_g . w)) with box bounds; w updated in place.
    Returns (loss, margins)."""
    n, p = Z.shape
    m = np.empty(n)
    for i in range(n):
        s = 0.0
        for q in range(p):
            s += Z[i, q] * w[q]
        m[i] = s
    E = np.empty(n)
    En = np.empty(n)
    cur = 0.0
    for i in range(n):
        lv, E[i] = _lrow_e(m[i])
        cur += c[i] * lv
    g = np.empty(p)
    H = np.empty((p, p))
    wn = np.empty(p)
    mn = np.empty(n)
    fi = np.empty(p, np.int64)
    d = np.empty(p)
    A = np.empty((p, p))
    bb = np.empty(p)
    Lm = np.empty((p, p))
    sol = np.empty(p)
    for _ in range(maxit):
        g[:] = 0.0
        H[:, :] = 0.0
        for i in range(n):
            pr = _sig_e(m[i], E[i])
            gi = -c[i] * pr
            hi_ = c[i] * pr * (1.0 - pr)
            for q in range(p):
                zq = Z[i, q]
                g[q] += gi * zq
                hz = hi_ * zq
                for r in range(q, p):
                    H[q, r] += hz * Z[i, r]
        nf = 0
        for q in range(p):
            d[q] = 0.0
            if not ((w[q] <= lo[q] + 1e-12 and g[q] > 0) or (w[q] >= hi[q] - 1e-12 and g[q] < 0)):
                fi[nf] = q
                nf += 1
        if nf > 0:
            for a in range(nf):
                bb[a] = -g[fi[a]]
                for b in range(nf):
                    qa, qb = fi[a], fi[b]
                    A[a, b] = H[min(qa, qb), max(qa, qb)]
                A[a, a] += 1e-10 * (1.0 + A[a, a])
            chol_solve_buf(A, bb, nf, Lm, sol)
            for a in range(nf):
                d[fi[a]] = sol[a]
        t = 1.0
        new = cur
        ok = False
        while t > 1e-8:
            for q in range(p):
                v = w[q] + t * d[q]
                wn[q] = min(max(v, lo[q]), hi[q])
            new = 0.0
            for i in range(n):
                s = 0.0
                for q in range(p):
                    s += Z[i, q] * wn[q]
                mn[i] = s
                lv, En[i] = _lrow_e(s)
                new += c[i] * lv
            if new <= cur:
                ok = True
                break
            t *= 0.5
        if not ok:
            break
        dec = cur - new
        w[:] = wn
        m, mn = mn, m
        E, En = En, E
        cur = new
        if dec <= tol * cur:
            break
    return cur, m
def newton_gains(Gm, h00, valid, bound)

gain[q, j]: loss decrease of one Newton step on a new coefficient for column j (box-clipped), with the intercept's curvature projected out.

Expand source code
@_kernel()
def newton_gains(Gm, h00, valid, bound):
    """gain[q, j]: loss decrease of one Newton step on a new coefficient for column j (box-clipped), with the
    intercept's curvature projected out."""
    d = Gm.shape[0]
    P = h00.shape[0]
    GA = np.full((P, d), -1.0)
    BA = np.zeros((P, d))
    for j in range(d):
        if not valid[j]:
            continue
        for q in range(P):
            g = abs(Gm[j, q])
            h0 = Gm[j, P + q]
            heff = max(Gm[j, 2 * P + q] - h0 * h0 / max(h00[q], 1e-300), 1e-12)
            if g <= bound * heff:
                GA[q, j] = 0.5 * g * g / heff
                BA[q, j] = Gm[j, q] / heff
            else:
                GA[q, j] = bound * g - 0.5 * heff * bound * bound
                BA[q, j] = bound if Gm[j, q] > 0 else -bound
    return GA, BA
def newton_split(Zp, gy, a0, a1, th, bound, maxit, tol)

newton_fit for a child = parent support + one binary column, on the parent's groups: group g has a cell with x_j = 0 (weight a0[g]) and one with x_j = 1 (weight a1[g]) that share the parent part of the design, so the parent block of the gradient / Hessian and the margins are accumulated once per group. th = (b, parent points, new point), updated in place. Returns the loss.

Expand source code
@_kernel()
def newton_split(Zp, gy, a0, a1, th, bound, maxit, tol):
    """newton_fit for a child = parent support + one binary column, on the parent's groups: group g has a cell
    with x_j = 0 (weight a0[g]) and one with x_j = 1 (weight a1[g]) that share the parent part of the design, so
    the parent block of the gradient / Hessian and the margins are accumulated once per group. th = (b, parent
    points, new point), updated in place. Returns the loss."""
    ng, ps = Zp.shape
    p = ps + 2
    lo = np.full(p, -bound)
    hi = np.full(p, bound)
    lo[0] = -1e300
    hi[0] = 1e300
    m0 = np.empty(ng)
    E0 = np.empty(ng)
    E1 = np.empty(ng)
    F0 = np.empty(ng)
    F1 = np.empty(ng)
    n0 = np.empty(ng)
    cur = 0.0
    for g in range(ng):
        s = gy[g] * th[0]
        for r in range(ps):
            s += Zp[g, r] * th[r + 1]
        m0[g] = s
        if a0[g] > 0.0:
            lv, E0[g] = _lrow_e(s)
            cur += a0[g] * lv
        if a1[g] > 0.0:
            lv, E1[g] = _lrow_e(s + gy[g] * th[p - 1])
            cur += a1[g] * lv
    gr = np.empty(p)
    H = np.empty((p, p))
    thn = np.empty(p)
    fi = np.empty(p, np.int64)
    d = np.empty(p)
    A = np.empty((p, p))
    bb = np.empty(p)
    Lm = np.empty((p, p))
    sol = np.empty(p)
    zp = np.empty(ps + 1)
    for _ in range(maxit):
        gr[:] = 0.0
        H[:, :] = 0.0
        for g in range(ng):
            yg = gy[g]
            gs = 0.0
            hs = 0.0
            g1 = 0.0
            h1 = 0.0
            if a0[g] > 0.0:
                pr = _sig_e(m0[g], E0[g])
                gs -= a0[g] * pr
                hs += a0[g] * pr * (1.0 - pr)
            if a1[g] > 0.0:
                m1 = m0[g] + yg * th[p - 1]
                pr = _sig_e(m1, E1[g])
                g1 = -a1[g] * pr
                h1 = a1[g] * pr * (1.0 - pr)
                gs += g1
                hs += h1
            zp[0] = yg
            for r in range(ps):
                zp[r + 1] = Zp[g, r]
            for q in range(ps + 1):
                zq = zp[q]
                gr[q] += gs * zq
                hz = hs * zq
                for r in range(q, ps + 1):
                    H[q, r] += hz * zp[r]
                H[q, p - 1] += h1 * zq * yg
            gr[p - 1] += g1 * yg
            H[p - 1, p - 1] += h1 * yg * yg
        nf = 0
        for q in range(p):
            d[q] = 0.0
            if not ((th[q] <= lo[q] + 1e-12 and gr[q] > 0) or (th[q] >= hi[q] - 1e-12 and gr[q] < 0)):
                fi[nf] = q
                nf += 1
        if nf > 0:
            for a in range(nf):
                bb[a] = -gr[fi[a]]
                for b in range(nf):
                    qa, qb = fi[a], fi[b]
                    A[a, b] = H[min(qa, qb), max(qa, qb)]
                A[a, a] += 1e-10 * (1.0 + A[a, a])
            chol_solve_buf(A, bb, nf, Lm, sol)
            for a in range(nf):
                d[fi[a]] = sol[a]
        t = 1.0
        new = cur
        ok = False
        while t > 1e-8:
            for q in range(p):
                v = th[q] + t * d[q]
                thn[q] = min(max(v, lo[q]), hi[q])
            new = 0.0
            for g in range(ng):
                s = gy[g] * thn[0]
                for r in range(ps):
                    s += Zp[g, r] * thn[r + 1]
                n0[g] = s
                if a0[g] > 0.0:
                    lv, F0[g] = _lrow_e(s)
                    new += a0[g] * lv
                if a1[g] > 0.0:
                    lv, F1[g] = _lrow_e(s + gy[g] * thn[p - 1])
                    new += a1[g] * lv
            if new <= cur:
                ok = True
                break
            t *= 0.5
        if not ok:
            break
        dec = cur - new
        th[:] = thn
        m0, n0 = n0, m0
        E0, F0 = F0, E0
        E1, F1 = F1, E1
        cur = new
        if dec <= tol * cur:
            break
    return cur
def pick_per_var(g, var_of, V, m)

Columns with the largest g > 0, at most one per variable (its best), best first.

Expand source code
@_kernel()
def pick_per_var(g, var_of, V, m):
    """Columns with the largest g > 0, at most one per variable (its best), best first."""
    bestj = np.full(V, -1, np.int64)
    for j in range(g.shape[0]):
        if g[j] > 0:
            v = var_of[j]
            if bestj[v] < 0 or g[j] > g[bestj[v]]:
                bestj[v] = j
    nv = 0
    for v in range(V):
        if bestj[v] >= 0:
            nv += 1
    cand = np.empty(nv, np.int64)
    vals = np.empty(nv)
    u = 0
    for v in range(V):
        if bestj[v] >= 0:
            cand[u] = bestj[v]
            vals[u] = -g[bestj[v]]
            u += 1
    o = np.argsort(vals, kind="mergesort")
    return cand[o[:m]]
def refit_round(w, QT, var_of, lev_of, kind, ncode, tptr, tabv, y, c, bound, tol, n_mult)

Continuous logistic fit (box [-bound, bound]) on the support of w, rounded by the calibrated loss over a grid of scales (as the beam's final nodes are): a move that changes every point at once.

Expand source code
@_kernel()
def refit_round(w, QT, var_of, lev_of, kind, ncode, tptr, tabv, y, c, bound, tol, n_mult):
    """Continuous logistic fit (box [-bound, bound]) on the support of w, rounded by the calibrated loss over a
    grid of scales (as the beam's final nodes are): a move that changes every point at once."""
    d = w.shape[0]
    n = QT.shape[1]
    cnt = 0
    for j in range(d):
        if w[j] != 0.0:
            cnt += 1
    S = np.empty(cnt, np.int64)
    u = 0
    for j in range(d):
        if w[j] != 0.0:
            S[u] = j
            u += 1
    ginv = np.zeros(n, np.int64)
    code = np.empty(n, np.int64)
    for i in range(n):
        code[i] = 1 if y[i] > 0 else 0
    ginv, ng, gy, gc, rep = regroup(ginv, 1, code, 2, y, c)
    for q in range(cnt):
        j = S[q]
        v = var_of[j]
        if kind[v] == 0:
            l = lev_of[j]
            for i in range(n):
                code[i] = 1 if QT[v, i] >= l else 0
            nc = 2
        else:
            for i in range(n):
                code[i] = QT[v, i]
            nc = ncode[v]
        ginv, ng, gy, gc, rep = regroup(ginv, ng, code, nc, y, c)
    Z = group_design(QT, var_of, tptr, tabv, gy, rep, S)
    npos = 0.0
    tot = 0.0
    for g in range(ng):
        tot += gc[g]
        if gy[g] > 0:
            npos += gc[g]
    beta = np.zeros(cnt + 1)
    beta[0] = np.log(npos / (tot - npos))
    lo = np.full(cnt + 1, -bound)
    hi = np.full(cnt + 1, bound)
    lo[0] = -1e300
    hi[0] = 1e300
    newton_fit(Z, gc, beta, lo, hi, 50, tol)
    XS = np.empty((ng, cnt))
    for g in range(ng):
        for q in range(cnt):
            XS[g, q] = Z[g, q + 1] * gy[g]
    r, l = calib_round_kernel(XS, gy, gc, beta[1:].copy(), n_mult, bound)
    out = np.zeros(d)
    for q in range(cnt):
        out[S[q]] = r[q]
    return out, l
def regroup(ginv, ng, code, ncode, y, c)

Refine row groups by a column's value code. Returns (inv, ng2, group y, group weight, representative row).

Expand source code
@_kernel()
def regroup(ginv, ng, code, ncode, y, c):
    """Refine row groups by a column's value code. Returns (inv, ng2, group y, group weight, representative row)."""
    n = ginv.shape[0]
    inv = np.empty(n, np.int64)
    M = ng * ncode
    cnt = 0
    if M <= 4 * n + 4096:
        table = np.full(M, -1, np.int64)
        for i in range(n):
            key = ginv[i] * ncode + code[i]
            t = table[key]
            if t < 0:
                t = cnt
                table[key] = t
                cnt += 1
            inv[i] = t
    else:
        keys = np.empty(n, np.int64)
        for i in range(n):
            keys[i] = ginv[i] * ncode + code[i]
        order = np.argsort(keys)
        last = -1
        for t in range(n):
            i = order[t]
            if t == 0 or keys[i] != last:
                cnt += 1
                last = keys[i]
            inv[i] = cnt - 1
    gy = np.empty(cnt)
    gc = np.zeros(cnt)
    rep = np.full(cnt, -1, np.int64)
    for i in range(n):
        g = inv[i]
        gc[g] += c[i]
        if rep[g] < 0:
            rep[g] = i
            gy[g] = y[i]
    return inv, cnt, gy, gc, rep
def removal_state(s, dj, y, c, a, b)

Score bins and statistics at fixed (a, b) after s -= dj, with the per-row screening derivatives.

Expand source code
@_kernel()
def removal_state(s, dj, y, c, a, b):
    """Score bins and statistics at fixed (a, b) after s -= dj, with the per-row screening derivatives."""
    n = s.shape[0]
    s2 = s - dj
    inv, sv, Wp, Wn = bin_scores(s2, y, c)
    m = sv.shape[0]
    lp = np.empty(m)
    ln_ = np.empty(m)
    pp = np.empty(m)
    pn = np.empty(m)
    wq = np.empty(m)
    T = bin_stats(sv, Wp, Wn, a, b, lp, ln_, pp, pn, wq)
    r = np.empty(n)
    h = np.empty(n)
    for i in range(n):
        q = inv[i]
        r[i] = (-pp[q] if y[i] > 0 else pn[q]) * c[i]
        h[i] = wq[q] * c[i]
    return s2, inv, sv, Wp, Wn, lp, ln_, pp, pn, wq, T, r, h
def round_all(GOFF, PNG, GY, GC, GREP, PS, plen, PW, QT, var_of, tptr, tabv, n_mult, bound, hv, N)

calib_round for every final beam node: rounded points (aligned with PS), calibrated loss / N, and a key of the score vector (hv . s) that identifies equivalent points.

Expand source code
@_kernel()
def round_all(GOFF, PNG, GY, GC, GREP, PS, plen, PW, QT, var_of, tptr, tabv, n_mult, bound, hv, N):
    """calib_round for every final beam node: rounded points (aligned with PS), calibrated loss / N, and a key of
    the score vector (hv . s) that identifies equivalent points."""
    P = PNG.shape[0]
    R = np.zeros((P, PS.shape[1]))
    L = np.empty(P)
    K = np.empty(P)
    n = QT.shape[1]
    for t in range(P):
        p = plen[t]
        a0 = GOFF[t]
        ng = PNG[t]
        S = PS[t, :p].copy()
        XS = np.empty((ng, p))
        for g in range(ng):
            for q in range(p):
                XS[g, q] = xval(QT, var_of, tptr, tabv, GREP[a0 + g], S[q])
        r, l = calib_round_kernel(XS, GY[a0:a0 + ng], GC[a0:a0 + ng], PW[t, 1:p + 1].copy(), n_mult, bound)
        R[t, :p] = r
        L[t] = l / N
        sc = score_kernel(QT, var_of, tptr, tabv, S, r)
        kk = 0.0
        for i in range(n):
            kk += hv[i] * sc[i]
        K[t] = kk
    return R, L, K
def row_hash(Q0, r, ys0)
Expand source code
@_kernel()
def row_hash(Q0, r, ys0):
    V, n = Q0.shape
    h = ys0 * r[V]
    for v in range(V):
        rv = r[v]
        for i in range(n):
            h[i] += rv * Q0[v, i]
    return h
def scan_columns(X)

Per column: count of nonzeros, whether a value other than 0 / 1 occurs; bitsets (word, column) of the nonzeros (counts by popcount of the bitsets).

Expand source code
@_kernel()
def scan_columns(X):
    """Per column: count of nonzeros, whether a value other than 0 / 1 occurs; bitsets (word, column) of the
    nonzeros (counts by popcount of the bitsets)."""
    n, d = X.shape
    nw = (n + 63) // 64
    B = np.zeros((nw, d), np.uint64)
    bad = np.zeros(d, np.bool_)
    for w in range(nw):
        Bw = B[w]
        for i in range(w * 64, min(n, w * 64 + 64)):
            sh = np.uint64(i & 63)
            row = X[i]
            for j in range(d):
                v = row[j]
                Bw[j] |= np.uint64(v != 0.0) << sh
                bad[j] |= (v != 0.0) & (v != 1.0)
    cnt = np.zeros(d, np.int64)
    for w in range(nw):
        for j in range(d):
            x = B[w, j]
            x = x - ((x >> np.uint64(1)) & np.uint64(0x5555555555555555))
            x = (x & np.uint64(0x3333333333333333)) + ((x >> np.uint64(2)) & np.uint64(0x3333333333333333))
            x = (x + (x >> np.uint64(4))) & np.uint64(0x0F0F0F0F0F0F0F0F)
            cnt[j] += np.int64((x * np.uint64(0x0101010101010101)) >> np.uint64(56))
    return ~bad, cnt, B
def score_kernel(QT, var_of, tptr, tabv, S, wS)
Expand source code
@_kernel()
def score_kernel(QT, var_of, tptr, tabv, S, wS):
    n = QT.shape[1]
    s = np.zeros(n)
    for t in range(S.shape[0]):
        j = S[t]
        v = var_of[j]
        base = tptr[j]
        for i in range(n):
            s[i] += wS[t] * tabv[base + QT[v, i]]
    return s
def screen_cols(G12, p, cols, a, dl, m)

The m columns (ascending) of cols with the best second-order score min_delta u G + u^2 H / 2 (u = a d), G and H in columns p and p + 1 of G12.

Expand source code
@_kernel()
def screen_cols(G12, p, cols, a, dl, m):
    """The m columns (ascending) of `cols` with the best second-order score min_delta u G + u^2 H / 2 (u = a d),
    G and H in columns p and p + 1 of G12."""
    nc = cols.shape[0]
    sc = np.empty(nc)
    vb = 0.0
    for r in range(dl.shape[0]):
        vb = max(vb, abs(dl[r]))
    for t in range(nc):
        sc[t] = _qbest(G12[cols[t], p], G12[cols[t], p + 1], a, vb)
    # partial selection of the m smallest (ties: lower index first)
    m = min(m, nc)
    sel = np.empty(m, np.int64)
    cnt = 0
    for t in range(nc):
        v = sc[t]
        if cnt < m:
            u = cnt
            cnt += 1
        elif v < sc[sel[m - 1]]:
            u = m - 1
        else:
            continue
        while u > 0 and sc[sel[u - 1]] > v:
            sel[u] = sel[u - 1]
            u -= 1
        sel[u] = t
    return cols[np.sort(sel)]
def screen_newton(G12, p, cols, m)

The m columns of cols with the largest continuous Newton gain G^2 / H (scale free: the refitted map can absorb any step size), ascending.

Expand source code
@_kernel()
def screen_newton(G12, p, cols, m):
    """The m columns of `cols` with the largest continuous Newton gain G^2 / H (scale free: the refitted map can
    absorb any step size), ascending."""
    nc = cols.shape[0]
    sc = np.empty(nc)
    for t in range(nc):
        g = G12[cols[t], p]
        h = G12[cols[t], p + 1]
        sc[t] = -g * g / h if h > 1e-300 else 0.0
    m = min(m, nc)
    sel = np.empty(m, np.int64)
    cnt = 0
    for t in range(nc):
        v = sc[t]
        if cnt < m:
            u = cnt
            cnt += 1
        elif v < sc[sel[m - 1]]:
            u = m - 1
        else:
            continue
        while u > 0 and sc[sel[u - 1]] > v:
            sel[u] = sel[u - 1]
            u -= 1
        sel[u] = t
    return cols[np.sort(sel)]
def slide_moves(w, S, free, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind, ncode, vrp, vrows, vcols, vcp, ne, tr_a, tr_b, radius, SHP, SHN)

Threshold slides: a support column of a chain moves to another threshold of its variable (at most radius levels away) with the same points. The score changes by +-w_j on the band of codes between the two levels, so every slide of a column is estimated from one (bin, code) histogram of its variable. Returns the ne best per support column (estimate, removed column, added column, value).

Expand source code
@_kernel()
def slide_moves(w, S, free, inv, sv, lp, ln_, pp, pn, wq, T, a, b, y, c, QT, var_of, lev_of, kind, ncode, vrp,
                vrows, vcols, vcp, ne, tr_a, tr_b, radius, SHP, SHN):
    """Threshold slides: a support column of a chain moves to another threshold of its variable (at most `radius`
    levels away) with the same points. The score changes by +-w_j on the band of codes between the two levels,
    so every slide of a column is estimated from one (bin, code) histogram of its variable. Returns the ne best
    per support column (estimate, removed column, added column, value)."""
    m = S.shape[0]
    nb_ = sv.shape[0]
    o_est = np.full(m * ne, np.inf)
    o_rj = np.full(m * ne, -1, np.int64)
    o_aj = np.full(m * ne, -1, np.int64)
    o_v = np.zeros(m * ne)
    cb = np.empty(nb_, np.int64)
    cx = np.empty(nb_)
    cwp = np.empty(nb_)
    cwn = np.empty(nb_)
    for q in range(m):
        j = S[q]
        v = var_of[j]
        nc = ncode[v]
        if kind[v] != 0 or nc <= 2:
            continue
        l = lev_of[j]
        lo_l = max(1, l - radius)
        hi_l = min(nc - 1, l + radius)
        ntar = 0
        tl = np.empty(hi_l - lo_l + 1, np.int64)
        for l2 in range(lo_l, hi_l + 1):
            j2 = vcols[vcp[v] + l2 - 1]
            if l2 != l and free[j2]:
                tl[ntar] = l2
                ntar += 1
        if ntar == 0:
            continue
        HP = SHP[q]
        HN = SHN[q]
        deltas = np.full((ntar, 1), w[j])
        tlo, TB = make_tb(sv, nb_, -abs(w[j]), abs(w[j]), a, b)
        out = np.empty((2, ntar, 1))
        for t in range(ntar):
            l2 = tl[t]
            lo2 = min(l, l2)
            hi2 = max(l, l2)
            sg = 1.0 if l2 < l else -1.0  # rows of the band gain (lower level) or lose the indicator
            nt = 0
            for bq in range(nb_):
                wp_ = HP[bq, lo2] - HP[bq, hi2]
                wn_ = HN[bq, lo2] - HN[bq, hi2]
                if wp_ > 0.0 or wn_ > 0.0:
                    cb[nt] = bq
                    cx[nt] = sg
                    cwp[nt] = wp_
                    cwn[nt] = wn_
                    nt += 1
            eval_cells(t, deltas, cb, cx, cwp, cwn, nt, sv, lp, ln_, pp, pn, wq, T, a, b, out, tr_a, tr_b, tlo, TB)
        fo = np.argsort(out[0, :, 0], kind="mergesort")
        for u in range(min(ne, ntar)):
            t = fo[u]
            o_est[q * ne + u] = out[0, t, 0]
            o_rj[q * ne + u] = j
            o_aj[q * ne + u] = vcols[vcp[v] + tl[t] - 1]
            o_v[q * ne + u] = w[j]
    return o_est, o_rj, o_aj, o_v
def solve(X, y, k, bound=5, time_limit=60.0, profile='decile')

Integer points for the columns of X (y in {0, 1}): at most k nonzero, each in [-bound, bound], minimising the calibrated log loss. profile is "decile" (about 9 thresholds per numeric column) or "fine" (about 99). Returns (points, calibrated mean log loss, seconds per stage: data, beam, rounding, local search, stopped_early).

Expand source code
def solve(X, y, k, bound=COEF_BOUND, time_limit=60.0, profile="decile"):
    """Integer points for the columns of ``X`` (y in {0, 1}): at most ``k`` nonzero, each in [-bound, bound],
    minimising the calibrated log loss. ``profile`` is "decile" (about 9 thresholds per numeric column) or "fine"
    (about 99). Returns (points, calibrated mean log loss, seconds per stage: data, beam, rounding, local search,
    stopped_early)."""
    if not HAVE_NUMBA:
        raise ImportError("the risk score solver needs numba (pip install numba)")
    if profile not in PROFILES:
        raise ValueError(f"profile must be one of {sorted(PROFILES)}, got {profile!r}")
    _warmup()
    search = _Search(int(k), int(bound), float(time_limit), PROFILES[profile])
    points, loss, seconds = search.fit(np.asarray(X, dtype=np.float64), np.asarray(y))
    return points, loss, seconds, search.stopped
def start_state(w, QT, var_of, tptr, tabv, y, c, N)

ScoreState of the points w in one kernel: scores, bins and the calibrated (L, a, b).

Expand source code
@_kernel()
def start_state(w, QT, var_of, tptr, tabv, y, c, N):
    """ScoreState of the points w in one kernel: scores, bins and the calibrated (L, a, b)."""
    d = w.shape[0]
    cnt = 0
    for j in range(d):
        if w[j] != 0.0:
            cnt += 1
    S = np.empty(cnt, np.int64)
    wS = np.empty(cnt)
    u = 0
    for j in range(d):
        if w[j] != 0.0:
            S[u] = j
            wS[u] = w[j]
            u += 1
    s = score_kernel(QT, var_of, tptr, tabv, S, wS)
    inv, sv, Wp, Wn = bin_scores(s, y, c)
    if sv.shape[0] < 2:
        npos = 0.0
        for i in range(y.shape[0]):
            if y[i] > 0:
                npos += c[i]
        L, a, b = calibrate_bins(sv, Wp, Wn, 0.0, np.log(npos / (N - npos)))
    else:
        sd = np.std(s)
        L, a, b = calibrate_bins(sv / sd, Wp, Wn, 0.0, 0.0)
        a = a / sd
    return s, L, a, b, inv, sv, Wp, Wn
def support_hist(S, inv, nb_, y, c, QT, var_of, kind, ncode, lev_of, radius)

For every support column of a chain: weights of y = +1 / -1 per (score bin, code of its variable), suffix- summed over codes (entry [q, b, l]: rows of bin b with code >= l), from one pass over the rows.

Expand source code
@_kernel()
def support_hist(S, inv, nb_, y, c, QT, var_of, kind, ncode, lev_of, radius):
    """For every support column of a chain: weights of y = +1 / -1 per (score bin, code of its variable), suffix-
    summed over codes (entry [q, b, l]: rows of bin b with code >= l), from one pass over the rows."""
    k = S.shape[0]
    n = inv.shape[0]
    vs = np.empty(k, np.int64)
    mc = 1
    for q in range(k):
        vs[q] = var_of[S[q]]
        if kind[vs[q]] == 0:
            mc = max(mc, ncode[vs[q]] + 1)
    HP = np.zeros((k, nb_, mc))
    HN = np.zeros((k, nb_, mc))
    for q in range(k):
        if kind[vs[q]] != 0:
            vs[q] = -1
    for i in range(n):
        bq = inv[i]
        ci = c[i]
        if y[i] > 0:
            for q in range(k):
                if vs[q] >= 0:
                    HP[q, bq, QT[vs[q], i]] += ci
        else:
            for q in range(k):
                if vs[q] >= 0:
                    HN[q, bq, QT[vs[q], i]] += ci
    for q in range(k):
        if vs[q] < 0:
            continue
        nc = ncode[vs[q]]
        # suffix sums down to the lowest level a value change or a slide reads
        lo_q = max(0, lev_of[S[q]] - radius)
        for bq in range(nb_):
            for qq in range(nc - 1, lo_q - 1, -1):
                HP[q, bq, qq] += HP[q, bq, qq + 1]
                HN[q, bq, qq] += HN[q, bq, qq + 1]
    return HP, HN
def swap_phase(s, w, S, free, y, c, a, b, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp, nzv, n_screen, ne, tr_a, tr_b, cinv, csv)

Swap candidates (remove S[q], add a column with a value): per removal, the n_screen best swap-in columns by a second-order screen at the removal state, evaluated by the calibrated-loss estimate; returns the ne best per removal as arrays (estimate, removed column, added column, value, loss at fixed (a, b)).

Expand source code
@_kernel()
def swap_phase(s, w, S, free, y, c, a, b, QT, var_of, lev_of, kind, ncode, tptr, tabv, vrp, vrows, vcols, vcp,
               nzv, n_screen, ne, tr_a, tr_b, cinv, csv):
    """Swap candidates (remove S[q], add a column with a value): per removal, the n_screen best swap-in columns by
    a second-order screen at the removal state, evaluated by the calibrated-loss estimate; returns the ne best per
    removal as arrays (estimate, removed column, added column, value, loss at fixed (a, b))."""
    n = s.shape[0]
    k = S.shape[0]
    d = w.shape[0]
    nv = nzv.shape[0]
    R = np.empty((n, 2 * k))
    INV = np.empty((k, n), np.int64)
    BIN = List()  # per removal: (sv, lp, ln, pp, pn, wq) stacked, and T
    TT = np.empty((k, 6))
    AB = np.empty((k, 2))
    nb0 = csv.shape[0]
    key = np.empty(n, np.int64)
    for q in range(k):
        j = S[q]
        v = var_of[j]
        base = tptr[j]
        if kind[v] == 0:
            # binary column: the removal state's bins are the current bins split by x_j (rows keyed by (bin, x))
            l = lev_of[j]
            pres = np.zeros(2 * nb0, np.bool_)
            Wp2k = np.zeros(2 * nb0)
            Wn2k = np.zeros(2 * nb0)
            for i in range(n):
                kk = 2 * cinv[i] + (1 if QT[v, i] >= l else 0)
                key[i] = kk
                pres[kk] = True
                if y[i] > 0:
                    Wp2k[kk] += c[i]
                else:
                    Wn2k[kk] += c[i]
            npres = 0
            for kk in range(2 * nb0):
                npres += pres[kk]
            ks = np.empty(npres, np.int64)
            ksc = np.empty(npres)
            u = 0
            for kk in range(2 * nb0):
                if pres[kk]:
                    ks[u] = kk
                    ksc[u] = csv[kk >> 1] - w[j] * (kk & 1)
                    u += 1
            o = np.argsort(ksc, kind="mergesort")
            newid = np.empty(2 * nb0, np.int64)
            svt = np.empty(npres)
            m2 = -1
            for t in range(npres):
                if t == 0 or ksc[o[t]] != svt[m2]:
                    m2 += 1
                    svt[m2] = ksc[o[t]]
                newid[ks[o[t]]] = m2
            m2 += 1
            sv = svt[:m2].copy()
            Wp = np.zeros(m2)
            Wn = np.zeros(m2)
            for kk in range(2 * nb0):
                if pres[kk]:
                    Wp[newid[kk]] += Wp2k[kk]
                    Wn[newid[kk]] += Wn2k[kk]
            lp = np.empty(m2)
            ln_ = np.empty(m2)
            pp = np.empty(m2)
            pn = np.empty(m2)
            wq = np.empty(m2)
            AB[q, 0] = a
            AB[q, 1] = b
            T = bin_stats(sv, Wp, Wn, a, b, lp, ln_, pp, pn, wq)
            inv = np.empty(n, np.int64)
            for i in range(n):
                g = newid[key[i]]
                inv[i] = g
                R[i, 2 * q] = (-pp[g] if y[i] > 0 else pn[g]) * c[i]
                R[i, 2 * q + 1] = wq[g] * c[i]
        else:
            dj = np.empty(n)
            for i in range(n):
                dj[i] = w[j] * tabv[base + QT[v, i]]
            s2, inv, sv, Wp, Wn, lp, ln_, pp, pn, wq, T, r_, h_ = removal_state(s, dj, y, c, a, b)
            AB[q, 0] = a
            AB[q, 1] = b
            R[:, 2 * q] = r_
            R[:, 2 * q + 1] = h_
        INV[q] = inv
        st6 = np.empty((6, sv.shape[0]))
        st6[0] = sv
        st6[1] = lp
        st6[2] = ln_
        st6[3] = pp
        st6[4] = pn
        st6[5] = wq
        BIN.append(st6)
        TT[q] = T
    pow2 = np.zeros(2 * k, np.bool_)
    for q in range(k):
        pow2[2 * q + 1] = True
    G12 = colsum(QT, R, kind, ncode, vcp, vcols, lev_of, tptr, tabv, pow2, vrp, vrows)
    o_est = np.full(k * ne, np.inf)
    o_fx = np.full(k * ne, np.inf)
    o_rj = np.full(k * ne, -1, np.int64)
    o_aj = np.full(k * ne, -1, np.int64)
    o_v = np.zeros(k * ne)
    for q in range(k):
        j = S[q]
        cnt = 0
        for t in range(d):
            if free[t]:
                cnt += 1
        if cnt == 0:
            continue
        nonSq = np.empty(cnt, np.int64)
        cnt = 0
        for t in range(d):
            if free[t]:
                nonSq[cnt] = t
                cnt += 1
        aq = AB[q, 0]
        bq = AB[q, 1]
        if abs(a) * (csv[csv.shape[0] - 1] - csv[0]) >= SWNG:
            cols = screen_newton(G12, 2 * q, nonSq, n_screen)  # integer steps overshoot at a steep map
        else:
            cols = screen_cols(G12, 2 * q, nonSq, aq, nzv, n_screen)
        m = cols.shape[0]
        st6 = BIN[q]
        inv = INV[q]
        sv = st6[0]
        lp = st6[1]
        ln_ = st6[2]
        pp = st6[3]
        pn = st6[4]
        wq = st6[5]
        T = TT[q]
        vo = np.empty(m, np.int64)
        for t in range(m):
            vo[t] = var_of[cols[t]]
        o = np.argsort(vo, kind="mergesort")
        deltas = np.empty((m, nv))
        for t in range(m):
            deltas[t, :] = nzv
        out = np.empty((2, m, nv))
        eval_vars(cols[o], o, deltas, inv, sv.shape[0], sv, lp, ln_, pp, pn, wq, T, aq, bq, y, c, QT, var_of,
                  lev_of, kind, ncode, tptr, tabv, out, tr_a, tr_b, vrp, vrows)
        flat = out[0].ravel()
        fo = np.argsort(flat, kind="mergesort")
        u = 0
        for t in range(fo.shape[0]):
            if u >= ne:
                break
            f = fo[t]
            if not np.isfinite(flat[f]):
                continue
            r = f % nv
            qq = f // nv
            o_est[q * ne + u] = flat[f]
            o_fx[q * ne + u] = out[1].ravel()[f]
            o_rj[q * ne + u] = j
            o_aj[q * ne + u] = cols[qq]
            o_v[q * ne + u] = nzv[r]
            u += 1
    return o_est, o_rj, o_aj, o_v, o_fx
def top_order(est, m)

Indices of the m smallest finite entries of est, ties in index order.

Expand source code
@_kernel()
def top_order(est, m):
    """Indices of the m smallest finite entries of est, ties in index order."""
    o = np.argsort(est, kind="mergesort")
    cnt = 0
    for t in range(o.shape[0]):
        if np.isfinite(est[o[t]]):
            cnt += 1
    cnt = min(cnt, m)
    return o[:cnt]
def total_loss(ym, c)
Expand source code
@_kernel()
def total_loss(ym, c):
    s = 0.0
    for i in range(ym.shape[0]):
        s += c[i] * _lrow(ym[i])
    return s
def unique_rows(h)

np.unique(h, return_index=True) with the counts: first occurrence of every distinct value, in value order, and the number of rows with that value. A stable LSD radix sort on the order-preserving bit pattern of h gives the same order as numpy's stable sort.

Expand source code
@_kernel()
def unique_rows(h):
    """np.unique(h, return_index=True) with the counts: first occurrence of every distinct value, in value order,
    and the number of rows with that value. A stable LSD radix sort on the order-preserving bit pattern of h
    gives the same order as numpy's stable sort."""
    n = h.shape[0]
    key = np.empty(n, np.uint64)
    hb = h.view(np.uint64)
    top = np.uint64(1) << np.uint64(63)
    for i in range(n):
        u = hb[i]
        key[i] = ~u if (u & top) else (u | top)  # IEEE order -> unsigned order
    o = np.arange(n)
    o2 = np.empty(n, np.int64)
    cnt = np.empty(257, np.int64)
    for sh in range(0, 64, 8):
        cnt[:] = 0
        for t in range(n):
            cnt[((key[o[t]] >> np.uint64(sh)) & np.uint64(255)) + 1] += 1
        if cnt[1:].max() == n:
            continue  # all keys share this byte
        for b in range(256):
            cnt[b + 1] += cnt[b]
        for t in range(n):
            bb = (key[o[t]] >> np.uint64(sh)) & np.uint64(255)
            o2[cnt[bb]] = o[t]
            cnt[bb] += 1
        o, o2 = o2, o
    first = np.empty(n, np.int64)
    cw = np.zeros(n)
    m = -1
    for t in range(n):
        i = o[t]
        if t == 0 or h[i] != h[o[t - 1]]:
            m += 1
            first[m] = i
        cw[m] += 1.0
    m += 1
    return first[:m].copy(), cw[:m].copy()
def xval(QT, var_of, tptr, tabv, i, j)
Expand source code
@_kernel()
def xval(QT, var_of, tptr, tabv, i, j):
    return tabv[tptr[j] + QT[var_of[j], i]]

Classes

class Data (X, y01, ms_sqrt=0.0)
Expand source code
class Data:
    def colsum(self, R, pow2=False):
        if np.ndim(pow2) == 0:
            pow2 = np.full(R.shape[1], bool(pow2))
        return colsum(self.QT, np.ascontiguousarray(R), self.kind, self.ncode, self.vcp, self.vcols, self.lev_of,
                      self.tptr, self.tabv, np.asarray(pow2, np.bool_), self.vrp, self.vrows)

    def __init__(self, X, y01, ms_sqrt=0.0):
        X = np.ascontiguousarray(X, dtype=np.float64)
        n0, d = X.shape
        self.d = d
        ys0 = np.where(np.asarray(y01) > 0, 1.0, -1.0)
        isbin, cnt, B = scan_columns(X)
        const = np.where(isbin, (cnt == 0) | (cnt == n0), False)
        nb_cols = np.flatnonzero(~isbin)
        if len(nb_cols):
            const[nb_cols] = np.ptp(X[:, nb_cols], axis=0) == 0
        chainable = isbin & ~const
        cand = np.flatnonzero(chainable)
        order = cand[np.lexsort((cand, -cnt[cand]))]
        chain, lev, nch = build_chains(B, cnt, order, d)
        others = np.flatnonzero(~chainable)
        V = nch + len(others)
        kind = np.zeros(V, np.int64)
        ncode = np.zeros(V, np.int64)
        var_of = chain.copy()
        lev_of = lev.copy()
        # codes in a byte when every variable is a chain of at most 255 columns (less memory in every row pass)
        small = len(others) == 0 and (lev.max() if len(lev) else 0) <= 255
        Q0 = np.zeros((V, n0), np.uint8 if small else np.int32)
        chain_codes(B, n0, chain, lev, nch, Q0)
        for ch in range(nch):
            ncode[ch] = 0
        _max_at(ncode, chain[cand], lev[cand])
        ncode[:nch] += 1
        tabs = [None] * d
        for t, j in enumerate(others):
            v = nch + t
            u, inv = np.unique(X[:, j], return_inverse=True)
            Q0[v] = inv.ravel()
            ncode[v] = len(u)
            kind[v] = 1 if len(u) <= RAW_DENSE else 2
            var_of[j] = v
            lev_of[j] = 0
            tabs[j] = u.astype(np.float64)
        tptr = np.zeros(d + 1, np.int64)
        np.cumsum(ncode[var_of], out=tptr[1:])
        tabv = chain_tables(tptr, lev_of, var_of, nch)
        for j in others:
            tabv[tptr[j]:tptr[j + 1]] = tabs[j]
        # unique (codes, y) rows with counts, via a random projection hash (deterministic seed)
        r = _normals(12345, V + 1)
        h = row_hash(Q0, r, ys0)
        first, self.c = unique_rows(h)
        self.QT, vrp_, vrows_ = gather_cols_nz(Q0, first, kind)
        if self.QT.dtype != np.uint8 and ncode.max() <= 256:
            self.QT = self.QT.astype(np.uint8)  # codes in a byte: less memory traffic in every row pass
        self.y = np.ascontiguousarray(ys0[first])
        self.n = self.QT.shape[1]
        self.N = float(self.c.sum())
        self.yc = self.y * self.c
        self.V, self.kind, self.ncode, self.var_of, self.lev_of = V, kind, ncode, var_of, lev_of
        self.tptr, self.tabv = tptr, tabv
        vorder = np.lexsort((lev_of, var_of))
        self.vcols = vorder.astype(np.int64)
        self.vcp = np.zeros(V + 1, np.int64)
        self.vcp[1:] = np.cumsum(np.bincount(var_of, minlength=V))
        self.nchains = nch
        self.allbin = len(others) == 0 or bool(np.all((tabv == 0.0) | (tabv == 1.0)))  # chain tables are 0 / 1
        # rows that can have a nonzero value per variable (code > 0 for a chain, every row otherwise)
        self.vrp, self.vrows = vrp_, vrows_
        cn = self.c / self.N
        if self.allbin:  # x^2 = x for every column: one column sum
            M = self.colsum(cn[:, None], np.array([False]))[:, 0].copy()
            M2 = M
        else:
            MM = self.colsum(np.stack([cn, cn], 1), np.array([False, True]))
            M, M2 = MM[:, 0].copy(), MM[:, 1].copy()
        var = M2 - M * M
        self.norm = np.sqrt(np.maximum(var, 0.0) * self.N)  # centred column norm, as in FasterRisk
        self.valid = self.norm > 1e-9
        # minimum support (profile "fine"): a binary column must hold at least ms_sqrt * sqrt(n) (weighted) rows
        # on each side
        self.cntw = M * self.N
        self.msup = ms_sqrt * np.sqrt(self.N)
        colbin = np.ones(d, np.bool_)
        if not self.allbin:
            colbin[others] = [np.all((tabs[j] == 0) | (tabs[j] == 1)) for j in others]
        self.valid &= ~colbin | ((self.cntw >= self.msup) & (self.N - self.cntw >= self.msup))
        self.bvalid = self.valid.copy()  # columns the beam may add
        # threshold cells for the beam's diversity rules: a chain column's cell is its variable and the decile of
        # its row fraction (CELLB buckets), so thresholds of one variable that split the rows alike share a cell
        frac = self.cntw / self.N
        bucket = np.where(kind[var_of] == 0, np.minimum((frac * CELLB).astype(np.int64), CELLB - 1), 0)
        _, pv_ = np.unique(var_of * CELLB + bucket, return_inverse=True)
        self.pvar = pv_.ravel().astype(np.int64)
        self.npvar = int(self.pvar.max()) + 1 if d else 0
        # complement pairs of binary columns share one beam hash (the same continuous fit up to the intercept)
        self.crep = complement_rep(B, n0, np.flatnonzero(isbin & ~const).astype(np.int64), d,
                                   _ints5(B.shape[0]).astype(np.uint64) * np.uint64(2) + np.uint64(1))

Methods

def colsum(self, R, pow2=False)
Expand source code
def colsum(self, R, pow2=False):
    if np.ndim(pow2) == 0:
        pow2 = np.full(R.shape[1], bool(pow2))
    return colsum(self.QT, np.ascontiguousarray(R), self.kind, self.ncode, self.vcp, self.vcols, self.lev_of,
                  self.tptr, self.tabv, np.asarray(pow2, np.bool_), self.vrp, self.vrows)
class ILS (D, k, bound)

Best-improvement local search over integer points, scored by the calibrated loss.

Expand source code
class ILS:
    """Best-improvement local search over integer points, scored by the calibrated loss."""

    def __init__(self, D, k, bound):
        self.D, self.k = D, k
        self.est_margin = 1e-4
        self.hw = _normals(11, D.d)  # hash of a point: its random projection
        self.vis = _new_dict_f64()
        self.cands = {}  # final point -> candidates of its failing swap phase
        self.allv = np.arange(-bound, bound + 1, dtype=np.float64)
        self.nzv = self.allv[self.allv != 0]

    def run(self, w, max_iter=100, gate=np.inf):
        D = self.D
        w = w.astype(np.float64).copy()
        st = start_state(w, D.QT, D.var_of, D.tptr, D.tabv, D.y, D.c, D.N)
        out = ils_kernel(w, st[0], st[1], st[2], st[3], st[4], st[5], st[6], st[7], self.vis, self.hw, self.k,
                         max_iter, True, float(gate), D.N, NEXACT, NSCREEN, ADDSCREEN, self.est_margin, D.y, D.c,
                         D.QT, D.var_of, D.lev_of, D.kind, D.ncode, D.tptr, D.tabv, D.vrp, D.vrows, D.vcols, D.vcp,
                         D.valid, self.allv, self.nzv, TR_A, TR_B)
        if out[2]:
            self.cands[out[1].tobytes()] = out[3:]
        return out[0], out[1]

Methods

def run(self, w, max_iter=100, gate=inf)
Expand source code
def run(self, w, max_iter=100, gate=np.inf):
    D = self.D
    w = w.astype(np.float64).copy()
    st = start_state(w, D.QT, D.var_of, D.tptr, D.tabv, D.y, D.c, D.N)
    out = ils_kernel(w, st[0], st[1], st[2], st[3], st[4], st[5], st[6], st[7], self.vis, self.hw, self.k,
                     max_iter, True, float(gate), D.N, NEXACT, NSCREEN, ADDSCREEN, self.est_margin, D.y, D.c,
                     D.QT, D.var_of, D.lev_of, D.kind, D.ncode, D.tptr, D.tabv, D.vrp, D.vrows, D.vcols, D.vcp,
                     D.valid, self.allv, self.nzv, TR_A, TR_B)
    if out[2]:
        self.cands[out[1].tobytes()] = out[3:]
    return out[0], out[1]