Skip to content

Solvers¤

solve_flux_bg_weighted_jax_nansafe

camino.solve_flux_bg_weighted_jax_nansafe(img, err, bad, m_unit, EPS=1e-12) ¤

Source code in camino.py
1435
1436
1437
1438
1439
1440
1441
1442
1443
1444
1445
1446
1447
1448
1449
1450
1451
1452
1453
1454
def solve_flux_bg_weighted_jax_nansafe(img, err, bad, m_unit, EPS=1e-12):
    img0 = jnp.nan_to_num(img, nan=0.0, posinf=0.0, neginf=0.0)
    err0 = jnp.nan_to_num(err, nan=jnp.inf, posinf=jnp.inf, neginf=jnp.inf)

    good = (~bad) & jnp.isfinite(err0) & (err0 > 0) & jnp.isfinite(img0)

    # avoid dividing on bad pixels
    err_safe = jnp.where(good, err0, 1.0)
    w = jnp.where(good, 1.0 / (jnp.maximum(err_safe, EPS) ** 2), 0.0)

    Smm = jnp.sum(w * m_unit * m_unit)
    Smy = jnp.sum(w * m_unit * img0)
    Smb = jnp.sum(w * m_unit)
    Sbb = jnp.sum(w)
    Sby = jnp.sum(w * img0)

    det = Smm * Sbb - Smb * Smb + EPS
    f_star = (Smy * Sbb - Smb * Sby) / det
    b_star = (Smm * Sby - Smb * Smy) / det
    return f_star, b_star
scale_poisson_no_bg

camino.scale_poisson_no_bg(img, psf_unit, bad) ¤

Source code in camino.py
1601
1602
1603
1604
1605
1606
def scale_poisson_no_bg(img, psf_unit, bad):
    m = jnp.where(bad, 0.0, psf_unit)
    d = jnp.where(bad, 0.0, img)
    num = jnp.sum(d)
    den = jnp.sum(m) + 1e-30
    return num / den
scale_ls_const_bg_unweighted

camino.scale_ls_const_bg_unweighted(img, psf_unit, bad) ¤

Source code in camino.py
1609
1610
1611
1612
1613
1614
1615
1616
1617
1618
1619
1620
def scale_ls_const_bg_unweighted(img, psf_unit, bad):
    M = jnp.where(bad, 0.0, psf_unit)
    D = jnp.where(bad, 0.0, img)
    Smm = jnp.sum(M * M)
    Sm1 = jnp.sum(M)
    S11 = jnp.sum(~bad).astype(M.dtype)
    Sdm = jnp.sum(D * M)
    Sd1 = jnp.sum(D)
    det = Smm * S11 - Sm1 * Sm1 + 1e-30
    f = (S11 * Sdm - Sm1 * Sd1) / det
    b = (-Sm1 * Sdm + Smm * Sd1) / det
    return f, b