Generation now matches offline tests: full frame-aligned context (len%6==0, cap 1020), not seq[-48:]
Browse files- daisychain.py +7 -1
daisychain.py
CHANGED
|
@@ -111,7 +111,13 @@ class DaisyChain:
|
|
| 111 |
homopolymers). Sampling (greedy=False) trades that for variety; repetition_penalty
|
| 112 |
then discourages the repeat loops these small specialists fall into."""
|
| 113 |
m = self.models[domain]
|
| 114 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
t = torch.tensor([ids], device=self.dev)
|
| 116 |
bases, emitted = [], []
|
| 117 |
while sum(len(b) for b in bases) < length:
|
|
|
|
| 111 |
homopolymers). Sampling (greedy=False) trades that for variety; repetition_penalty
|
| 112 |
then discourages the repeat loops these small specialists fall into."""
|
| 113 |
m = self.models[domain]
|
| 114 |
+
# frame-align + cap: our 6-mer tokens tile cleanly only when the context length is a
|
| 115 |
+
# multiple of 6 (else the model generates out of phase). Trim the leading remainder and
|
| 116 |
+
# cap to the context window — this is what makes generation match the offline tests.
|
| 117 |
+
p = clean(prompt) if prompt else ""
|
| 118 |
+
p = p[-1020:] # 1020 = 170*6, within the 1024-token window
|
| 119 |
+
p = p[len(p) % 6:] # drop leading remainder so the 6-mers align to the end
|
| 120 |
+
ids = [self.bos] + (self.tok.encode(p, add_special_tokens=False) if p else [])
|
| 121 |
t = torch.tensor([ids], device=self.dev)
|
| 122 |
bases, emitted = [], []
|
| 123 |
while sum(len(b) for b in bases) < length:
|