(q, scale, eps)
| 51 | return torch.stack((q, k, v), dim=2) |
| 52 | |
| 53 | def 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 | |
| 60 | class SpatialTransformerSimpleV2(nn.Module): |
| 61 | """ |