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

Method get_freqs

k_diffusion/models/axial_rope.py:99–105  ·  view source on GitHub ↗
(self, pos)

Source from the content-addressed store, hash-verified

97 return f"dim={dim}, n_heads={self.n_heads}, start_index={self.start_index}"
98
99 def get_freqs(self, pos):
100 if pos.shape[-1] != 2:
101 raise ValueError("input shape must be (..., 2)")
102 freqs_h = pos[..., None, None, 0] * self.freqs_h.exp()
103 freqs_w = pos[..., None, None, 1] * self.freqs_w.exp()
104 freqs = torch.cat((freqs_h, freqs_w), dim=-1).repeat_interleave(2, dim=-1)
105 return freqs.transpose(-2, -3)
106
107 def forward(self, x, pos):
108 freqs = self.get_freqs(pos)

Callers 1

forwardMethod · 0.95

Calls 1

transposeMethod · 0.80

Tested by

no test coverage detected