Cleanup old math versions and fix variables (#16)

This commit is contained in:
Reithan
2025-07-26 04:34:34 -07:00
committed by GitHub
parent 62bef2e275
commit 3d8827f132
+21 -191
View File
@@ -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: