Switch sampled decoder to nucleus/top-p (Carbon's decoder) instead of top-k — escapes low-complexity loops
Browse files- 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 136 |
logits[0, tid] = -1e9
|
| 137 |
ti = int(logits.argmax())
|
| 138 |
else:
|
|
|
|
|
|
|
|
|
|
| 139 |
logits = logits / max(temperature, 1e-6)
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
|
|
|
|
|
|
| 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])
|