Quazim0t0 commited on
Commit
c4d1e75
·
verified ·
1 Parent(s): 6fb2512

Switch sampled decoder to nucleus/top-p (Carbon's decoder) instead of top-k — escapes low-complexity loops

Browse files
Files changed (1) hide show
  1. daisychain.py +16 -10
daisychain.py CHANGED
@@ -103,7 +103,7 @@ class DaisyChain:
103
 
104
  @torch.no_grad()
105
  def generate_stream(self, domain, length=180, temperature=1.0, top_k=40,
106
- repetition_penalty=1.3, prompt="", greedy=False):
107
  """Yield the growing continuation base-by-base (for live streaming UIs).
108
  greedy=True takes the argmax 6-mer each step (deterministic — the model's single
109
  best guess, same decoding as the sequence-recovery metric; the coding domains
@@ -127,20 +127,26 @@ class DaisyChain:
127
  # self-reinforcing loop (argmax keeps re-picking the same 6-mer); the penalty plus
128
  # a hard block on the last few emitted tokens forces it onward to its next-best,
129
  # non-repeating guess — which is what actually de-degenerates the output.
130
- if repetition_penalty and repetition_penalty != 1.0:
131
- for tid in set(emitted[-12:]):
132
- v = logits[0, tid]
133
- logits[0, tid] = v / repetition_penalty if v > 0 else v * repetition_penalty
134
  if greedy:
135
- for tid in set(emitted[-8:]): # block longer low-complexity cycles
 
 
 
 
 
136
  logits[0, tid] = -1e9
137
  ti = int(logits.argmax())
138
  else:
 
 
 
139
  logits = logits / max(temperature, 1e-6)
140
- if top_k > 0:
141
- v, _ = torch.topk(logits, top_k)
142
- logits[logits < v[:, [-1]]] = -1e9
143
- ti = int(torch.multinomial(F.softmax(logits, dim=-1), 1))
 
 
144
  emitted.append(ti)
145
  t = torch.cat([t, torch.tensor([[ti]], device=self.dev)], dim=1)
146
  bases.append(self.tok._ids_to_tokens[ti])
 
103
 
104
  @torch.no_grad()
105
  def generate_stream(self, domain, length=180, temperature=1.0, top_k=40,
106
+ repetition_penalty=1.3, prompt="", greedy=False, top_p=0.9):
107
  """Yield the growing continuation base-by-base (for live streaming UIs).
108
  greedy=True takes the argmax 6-mer each step (deterministic — the model's single
109
  best guess, same decoding as the sequence-recovery metric; the coding domains
 
127
  # self-reinforcing loop (argmax keeps re-picking the same 6-mer); the penalty plus
128
  # a hard block on the last few emitted tokens forces it onward to its next-best,
129
  # non-repeating guess — which is what actually de-degenerates the output.
 
 
 
 
130
  if greedy:
131
+ # argmax with repeat penalty + a hard block on recent tokens (greedy alone loops)
132
+ if repetition_penalty and repetition_penalty != 1.0:
133
+ for tid in set(emitted[-12:]):
134
+ v = logits[0, tid]
135
+ logits[0, tid] = v / repetition_penalty if v > 0 else v * repetition_penalty
136
+ for tid in set(emitted[-8:]):
137
  logits[0, tid] = -1e9
138
  ti = int(logits.argmax())
139
  else:
140
+ # nucleus (top-p) sampling — the same decoder Carbon uses. Adapts the candidate
141
+ # set to the distribution (keeps only tokens carrying mass), which escapes the
142
+ # low-complexity GC/AT loops that a fixed top-k falls into. No repetition penalty.
143
  logits = logits / max(temperature, 1e-6)
144
+ sl, si = torch.sort(logits[0], descending=True)
145
+ cum = torch.cumsum(F.softmax(sl, dim=-1), dim=-1)
146
+ rm = cum > top_p
147
+ rm[1:] = rm[:-1].clone(); rm[0] = False
148
+ logits[0, si[rm]] = -1e9
149
+ ti = int(torch.multinomial(F.softmax(logits[0], dim=-1), 1))
150
  emitted.append(ti)
151
  t = torch.cat([t, torch.tensor([[ti]], device=self.dev)], dim=1)
152
  bases.append(self.tok._ids_to_tokens[ti])