ThingsAI commited on
Commit
4e1e026
·
verified ·
1 Parent(s): 176de1f

fix: RotaryEmbedding lazy cache build (evita garbage da meta-device init)

Browse files
Files changed (1) hide show
  1. modeling_quark.py +20 -10
modeling_quark.py CHANGED
@@ -28,14 +28,18 @@ class RotaryEmbedding(nn.Module):
28
  super().__init__()
29
  inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2).float() / head_dim))
30
  self.register_buffer("inv_freq", inv_freq, persistent=False)
31
- self._build_cache(max_seq_len)
32
-
33
- def _build_cache(self, seq_len):
34
- t = torch.arange(seq_len, device=self.inv_freq.device).float()
35
- freqs = torch.outer(t, self.inv_freq)
 
 
 
 
36
  emb = torch.cat([freqs, freqs], dim=-1)
37
- self.register_buffer("cos_cache", emb.cos()[None, None], persistent=False)
38
- self.register_buffer("sin_cache", emb.sin()[None, None], persistent=False)
39
  self._max = seq_len
40
 
41
  @staticmethod
@@ -45,11 +49,17 @@ class RotaryEmbedding(nn.Module):
45
 
46
  def forward(self, q, k):
47
  T = q.size(2)
48
- if T > self._max:
49
- self._build_cache(T)
 
 
 
 
 
 
 
50
  cos = self.cos_cache[:, :, :T, :]
51
  sin = self.sin_cache[:, :, :T, :]
52
- # Identico a train.py — nessun cast, broadcast naturale
53
  q = q * cos + self._rotate_half(q) * sin
54
  k = k * cos + self._rotate_half(k) * sin
55
  return q, k
 
28
  super().__init__()
29
  inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2).float() / head_dim))
30
  self.register_buffer("inv_freq", inv_freq, persistent=False)
31
+ self.max_seq_len = max_seq_len
32
+ self._max = 0 # forza build al primo forward, mai nell'__init__
33
+ self.cos_cache = None
34
+ self.sin_cache = None
35
+
36
+ def _build_cache(self, seq_len, device, dtype):
37
+ # Ricalcola sempre da inv_freq corrente (mai cache stantia da meta-device)
38
+ t = torch.arange(seq_len, device=device, dtype=torch.float32)
39
+ freqs = torch.outer(t, self.inv_freq.to(device=device, dtype=torch.float32))
40
  emb = torch.cat([freqs, freqs], dim=-1)
41
+ self.cos_cache = emb.cos()[None, None].to(dtype)
42
+ self.sin_cache = emb.sin()[None, None].to(dtype)
43
  self._max = seq_len
44
 
45
  @staticmethod
 
49
 
50
  def forward(self, q, k):
51
  T = q.size(2)
52
+ # Ricostruisce la cache se: mai costruita, troppo corta, o device/dtype cambiati
53
+ needs_rebuild = (
54
+ self.cos_cache is None
55
+ or T > self._max
56
+ or self.cos_cache.device != q.device
57
+ or self.cos_cache.dtype != q.dtype
58
+ )
59
+ if needs_rebuild:
60
+ self._build_cache(max(T, self.max_seq_len), q.device, q.dtype)
61
  cos = self.cos_cache[:, :, :T, :]
62
  sin = self.sin_cache[:, :, :T, :]
 
63
  q = q * cos + self._rotate_half(q) * sin
64
  k = k * cos + self._rotate_half(k) * sin
65
  return q, k