(self, ids: torch.Tensor)
| 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 | |
| 1317 | class TimestepEmbedding(nn.Module): |
nothing calls this directly
no test coverage detected