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