Cleanup old math versions and fix variables (#16)
This commit is contained in:
+21
-191
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user