(self, dim, max_position_embeddings=2048, base=10000, device=None)
| 21 | |
| 22 | class QuantLlamaRotaryEmbedding(nn.Module): |
| 23 | def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None): |
| 24 | super().__init__() |
| 25 | |
| 26 | self.dim = dim |
| 27 | self.max_position_embeddings = max_position_embeddings |
| 28 | self.base = base |
| 29 | inv_freq = 1.0 / ( |
| 30 | self.base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim) |
| 31 | ) |
| 32 | self.register_buffer("inv_freq", inv_freq) |
| 33 | # Build here to make `torch.jit.trace` work. |
| 34 | self._set_cos_sin_cache( |
| 35 | seq_len=max_position_embeddings, |
| 36 | device=self.inv_freq.device, |
| 37 | dtype=torch.get_default_dtype(), |
| 38 | ) |
| 39 | |
| 40 | def _set_cos_sin_cache(self, seq_len, device, dtype): |
| 41 | self.max_seq_len_cached = seq_len |
nothing calls this directly
no test coverage detected