(qkv, scale, eps)
| 46 | return q * scale_q.to(q.dtype), k * scale_k.to(k.dtype) |
| 47 | |
| 48 | def scale_for_cosine_sim_qkv(qkv, scale, eps): |
| 49 | q, k, v = qkv.unbind(2) |
| 50 | q, k = scale_for_cosine_sim(q, k, scale[:, None], 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)) |
no test coverage detected