From 3d8827f13245ab0110fa9a87d2881eb43055ce9b Mon Sep 17 00:00:00 2001 From: Reithan Date: Sat, 26 Jul 2025 04:34:34 -0700 Subject: [PATCH] Cleanup old math versions and fix variables (#16) --- NRS/nodes_NRS.py | 212 +++++------------------------------------------ 1 file changed, 21 insertions(+), 191 deletions(-) diff --git a/NRS/nodes_NRS.py b/NRS/nodes_NRS.py index ffc440a..5986995 100644 --- a/NRS/nodes_NRS.py +++ b/NRS/nodes_NRS.py @@ -161,7 +161,7 @@ class NRS: sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1)) sig_root = (sigma ** 2 + 1).sqrt() - nrs_cond, nrs_uncond = None, None + x_div, nrs_cond, nrs_uncond = None, None, None match self.__OPERATION_SPACE: case PredictionType.V: x_div, nrs_cond, nrs_uncond = self._convert_to_v_space(x_orig, sig_root, sigma, cond, uncond) @@ -174,201 +174,31 @@ class NRS: case _: raise RuntimeError("NRS.nrs: Invalid PredictionType used.") - x_final = None - match "v0.6.0": - case "v1": - # displace cond by rejection of uncond on cond - u_dot_c = torch.sum(nrs_uncond * nrs_cond, dim=-1, keepdim=True) - c_dot_c = torch.sum(nrs_cond * nrs_cond, dim=-1, keepdim=True) - u_on_c = (u_dot_c / c_dot_c) * nrs_cond - u_rej_c = nrs_uncond - u_on_c - displaced = (nrs_cond - skew * u_rej_c) - logging.debug(f"NRS.nrs: displaced") + def _dot(a, b): + return (a*b).sum(dim=1, keepdim=True) # [B,C,W,H] => [B,1,W,H] - # squash displaced vector towards len(cond) based on squash scale - d_len_sq = torch.sum(displaced * displaced, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * ((c_dot_c/d_len_sq) ** 0.5) - squashed = displaced * squash_scale - logging.debug(f"NRS.nrs: squashed") + def _nrm2(v): + return _dot(v, v) - # stretch turned vector towards cond based on stretch scale - sq_dot_c = torch.sum(squashed * nrs_cond, dim=-1, keepdim=True) - sq_on_c = (sq_dot_c / c_dot_c) * nrs_cond - x_final = squashed + sq_on_c * stretch - logging.debug(f"NRS.nrs: final") - case "v2": - # displace cond by rejection of uncond on cond - u_dot_c = torch.sum(nrs_uncond * nrs_cond, dim=-1, keepdim=True) - c_dot_c = torch.sum(nrs_cond * nrs_cond, dim=-1, keepdim=True) - u_on_c_mag = (u_dot_c / c_dot_c) - u_on_c = u_on_c_mag * nrs_cond - u_rej_c = nrs_uncond - u_on_c - displaced = nrs_cond + stretch * (nrs_cond - torch.clamp(u_dot_c / c_dot_c, min=0, max=1) * nrs_cond) - skew * u_rej_c - logging.debug(f"NRS.nrs: displaced & stretched") + eps = torch.finfo(nrs_cond.dtype).eps + c_dot_c = _nrm2(nrs_cond) + eps # [B,1,W,H] + u_dot_c = _dot(nrs_uncond, nrs_cond) # [B,1,W,H] + u_on_c = (u_dot_c / c_dot_c) * nrs_cond # [B,1,W,H] * [B,C,H,W] + + # Amplify Cond based on length compared to projection of uncond + proj_diff = nrs_cond - u_on_c + stretched = nrs_cond + (stretch * proj_diff) - # squash displaced vector towards len(cond) based on squash scale - d_len_sq = torch.sum(displaced * displaced, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * ((c_dot_c/d_len_sq) ** 0.5) - x_final = displaced * squash_scale - logging.debug(f"NRS.nrs: final") - case "v3": - # displace cond by rejection of uncond on cond - u_dot_c = torch.sum(nrs_uncond * nrs_cond, dim=-1, keepdim=True) - c_dot_c = torch.sum(nrs_cond * nrs_cond, dim=-1, keepdim=True) - u_on_c_mag = (u_dot_c / c_dot_c) - u_on_c = u_on_c_mag * nrs_cond - u_rej_c = nrs_uncond - u_on_c - displaced = (nrs_cond - skew * u_rej_c) - logging.debug(f"NRS.nrs: displaced") + # Skew/Steer Conf based on rejection of uncond on cond + u_rej_c = nrs_uncond - u_on_c + skewed = stretched - (skew * u_rej_c) - # squash displaced vector towards len(cond) based on squash scale - d_len_sq = torch.sum(displaced * displaced, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * ((c_dot_c/d_len_sq) ** 0.5) + # Squash final length back down to original length of cond + cond_len = nrs_cond.norm(dim=1, keepdim=True) + nrs_len = skewed.norm(dim=1, keepdim=True) + eps - # stretch vector towards 2*len(cond) - len(u_on_c) - c_len = c_dot_c ** 0.5 - stretch_scale = (1 - stretch) + stretch * (2 * c_len - u_on_c_mag)/c_len - - x_final = displaced * squash_scale * stretch_scale - logging.debug(f"NRS.nrs: final") - case "v4": - u_dot_c = torch.sum(nrs_uncond * nrs_cond, dim=-1, keepdim=True) - c_dot_c = torch.sum(nrs_cond * nrs_cond, dim=-1, keepdim=True) - u_on_c_mag = (u_dot_c / c_dot_c) - u_on_c = u_on_c_mag * nrs_cond - u_rej_c = nrs_uncond - u_on_c - rej_dor_rej = torch.sum(u_rej_c * u_rej_c, dim=-1, keepdim=True) - x_final = (nrs_cond - squash * u_rej_c + stretch * nrs_cond * ((rej_dor_rej/c_dot_c) ** 0.5)) - logging.debug(f"NRS.nrs: displaced") - case "v0.4.1": - u_dot_c = torch.sum(nrs_uncond * nrs_cond, dim=-1, keepdim=True) - c_dot_c = torch.sum(nrs_cond * nrs_cond, dim=-1, keepdim=True) - u_on_c_mag = (u_dot_c / c_dot_c) - u_on_c = u_on_c_mag * nrs_cond - u_rej_c = nrs_uncond - u_on_c - rej_dor_rej = torch.sum(u_rej_c * u_rej_c, dim=-1, keepdim=True) - stretched = nrs_cond + stretch * nrs_cond * ((rej_dor_rej/c_dot_c) ** 0.5) - skewed = stretched - skew * u_rej_c - sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * ((c_dot_c/sk_dot_sk) ** 0.5) - x_final = skewed * squash_scale - logging.debug(f"NRS.nrs: displaced") - case "v0.4.2": - u_dot_c = torch.sum(nrs_uncond * nrs_cond, dim=-1, keepdim=True) - c_dot_c = torch.sum(nrs_cond * nrs_cond, dim=-1, keepdim=True) - u_on_c_mag = (u_dot_c / c_dot_c) - u_on_c = u_on_c_mag * nrs_cond - u_rej_c = nrs_uncond - u_on_c - proj_len = torch.sum(u_on_c * u_on_c, dim=-1, keepdim=True) ** 0.5 - cond_len = c_dot_c ** 0.5 - stretched = nrs_cond * (1 + stretch * torch.abs(cond_len - proj_len) / cond_len) - skewed = stretched - skew * u_rej_c - sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * cond_len / (sk_dot_sk ** 0.5) - x_final = skewed * squash_scale - logging.debug(f"NRS.nrs: displaced") - case "v0.4.3": - u_dot_c = torch.sum(nrs_uncond * nrs_cond, dim=-1, keepdim=True) - c_dot_c = torch.sum(nrs_cond * nrs_cond, dim=-1, keepdim=True) - u_on_c_mag = (u_dot_c / c_dot_c) - u_on_c = u_on_c_mag * nrs_cond - u_rej_c = nrs_uncond - u_on_c - proj_len = torch.sum(u_on_c * u_on_c, dim=-1, keepdim=True) ** 0.5 - cond_len = c_dot_c ** 0.5 - stretched = nrs_cond * (1 + stretch * (cond_len - proj_len) / cond_len) - skewed = stretched - skew * u_rej_c - sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * cond_len / (sk_dot_sk ** 0.5) - x_final = skewed * squash_scale - logging.debug(f"NRS.nrs: displaced") - case "v0.4.4": - u_dot_c = torch.sum(nrs_uncond * nrs_cond, dim=-1, keepdim=True) - c_dot_c = torch.sum(nrs_cond * nrs_cond, dim=-1, keepdim=True) - u_on_c_mag = (u_dot_c / c_dot_c) - u_on_c = u_on_c_mag * nrs_cond - u_rej_c = nrs_uncond - u_on_c - cond_len = c_dot_c ** 0.5 - proj_diff = nrs_cond - u_on_c - proj_diff_len = torch.sum(proj_diff * proj_diff, dim=-1, keepdim=True) ** 0.5 - stretched = nrs_cond * (1 + stretch * proj_diff_len / cond_len) - skewed = stretched - skew * u_rej_c - sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * cond_len / (sk_dot_sk ** 0.5) - x_final = skewed * squash_scale - logging.debug(f"NRS.nrs: displaced") - case "v0.4.5": - u_dot_c = torch.sum(nrs_uncond * nrs_cond, dim=-1, keepdim=True) - c_dot_c = torch.sum(nrs_cond * nrs_cond, dim=-1, keepdim=True) - u_on_c_mag = (u_dot_c / c_dot_c) - u_on_c = u_on_c_mag * nrs_cond - u_rej_c = nrs_uncond - u_on_c - cond_len = c_dot_c ** 0.5 - proj_diff = nrs_cond - u_on_c - - # Amplify Cond based on length compared to projection of uncond - stretched = nrs_cond + (stretch * proj_diff) - - # Skew/Steer Conf based on rejection of uncond on cond - skewed = stretched - skew * u_rej_c - - # Squash final length back down to original length of cond - sk_dot_sk = torch.sum(skewed * skewed, dim=-1, keepdim=True) - squash_scale = (1 - squash) + squash * cond_len / (sk_dot_sk ** 0.5) - x_final = skewed * squash_scale - case "v0.5.0": - def _dot(a, b): - return (a*b).flatten(2).sum(dim=2, keepdim=True) # [B,C,W,H] => [B,C,1] - - def _nrm2(v): - return _dot(v, v) - - eps = torch.finfo(nrs_cond.dtype).eps - c_dot_c = _nrm2(nrs_cond) + eps # [B,1] - u_dot_c = _dot(nrs_uncond, nrs_cond) # [B,1] - - u_on_c = (u_dot_c / c_dot_c).unsqueeze(-1) * nrs_cond # [B,1,1,1] * [B,C,H,W] - - # Amplify Cond based on length compared to projection of uncond - proj_diff = nrs_cond - u_on_c - stretched = nrs_cond + (stretch * proj_diff) - - # Skew/Steer Conf based on rejection of uncond on cond - u_rej_c = nrs_uncond - u_on_c - skewed = stretched - (skew * u_rej_c) - - # Squash final length back down to original length of cond - cond_len = torch.sqrt(c_dot_c) # [B,1] - nrs_len = torch.sqrt(_nrm2(skewed)) + eps # [B,1] - - squash_scale = (1 - squash) + (squash * (cond_len / nrs_len)) - x_final = skewed * squash_scale.unsqueeze(-1) - case "v0.6.0": - def _dot(a, b): - return (a*b).sum(dim=1, keepdim=True) # [B,C,W,H] => [B,1,W,H] - - def _nrm2(v): - return _dot(v, v) - - eps = torch.finfo(nrs_cond.dtype).eps - c_dot_c = _nrm2(nrs_cond) + eps # [B,1] - u_dot_c = _dot(nrs_uncond, nrs_cond) # [B,1] - - u_on_c = (u_dot_c / c_dot_c) * nrs_cond # [B,1,1,1] * [B,C,H,W] - - # Amplify Cond based on length compared to projection of uncond - proj_diff = nrs_cond - u_on_c - stretched = nrs_cond + (stretch * proj_diff) - - # Skew/Steer Conf based on rejection of uncond on cond - u_rej_c = nrs_uncond - u_on_c - skewed = stretched - (skew * u_rej_c) - - # Squash final length back down to original length of cond - cond_len = cond.norm(dim=1, keepdim=True) - nrs_len = skewed.norm(dim=1, keepdim=True) - - squash_scale = (1 - squash) + (squash * (cond_len / nrs_len)) - x_final = skewed * squash_scale + squash_scale = (1 - squash) + (squash * (cond_len / nrs_len)) + x_final = skewed * squash_scale match self.__OPERATION_SPACE: case PredictionType.V: