MCPcopy Create free account
hub / github.com/Meshcapade/difflocks / scale_for_cosine_sim_single

Function scale_for_cosine_sim_single

k_diffusion/models/attention.py:53–58  ·  view source on GitHub ↗
(q, scale, eps)

Source from the content-addressed store, hash-verified

51 return torch.stack((q, k, v), dim=2)
52
53def scale_for_cosine_sim_single(q, scale, eps):
54 dtype = reduce(torch.promote_types, (q.dtype, scale.dtype, torch.float32))
55 sum_sq_q = torch.sum(q.to(dtype)**2, dim=-1, keepdim=True)
56 sqrt_scale = torch.sqrt(scale.to(dtype))
57 scale_q = sqrt_scale * torch.rsqrt(sum_sq_q + eps)
58 return q * scale_q.to(q.dtype)
59
60class SpatialTransformerSimpleV2(nn.Module):
61 """

Callers 1

forwardMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected