MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / __init__

Method __init__

inference/modules/fused_attn.py:23–38  ·  view source on GitHub ↗
(self, dim, max_position_embeddings=2048, base=10000, device=None)

Source from the content-addressed store, hash-verified

21
22class 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

Callers

nothing calls this directly

Calls 2

_set_cos_sin_cacheMethod · 0.95
__init__Method · 0.45

Tested by

no test coverage detected