MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / forward

Method forward

architecture/embeddings.py:1293–1314  ·  view source on GitHub ↗
(self, ids: torch.Tensor)

Source from the content-addressed store, hash-verified

1291 self.axes_dim = axes_dim
1292
1293 def forward(self, ids: torch.Tensor) -> torch.Tensor:
1294 n_axes = ids.shape[-1]
1295 cos_out = []
1296 sin_out = []
1297 pos = ids.float()
1298 is_mps = ids.device.type == "mps"
1299 is_npu = ids.device.type == "npu"
1300 freqs_dtype = torch.float32 if (is_mps or is_npu) else torch.float64
1301 for i in range(n_axes):
1302 cos, sin = get_1d_rotary_pos_embed(
1303 self.axes_dim[i],
1304 pos[:, i],
1305 theta=self.theta,
1306 repeat_interleave_real=True,
1307 use_real=True,
1308 freqs_dtype=freqs_dtype,
1309 )
1310 cos_out.append(cos)
1311 sin_out.append(sin)
1312 freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device)
1313 freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device)
1314 return freqs_cos, freqs_sin
1315
1316
1317class TimestepEmbedding(nn.Module):

Callers

nothing calls this directly

Calls 2

get_1d_rotary_pos_embedFunction · 0.85
toMethod · 0.45

Tested by

no test coverage detected