Upload daisychain.py with huggingface_hub
Browse files- daisychain.py +17 -8
daisychain.py
CHANGED
|
@@ -102,25 +102,34 @@ class DaisyChain:
|
|
| 102 |
return best, bpb
|
| 103 |
|
| 104 |
@torch.no_grad()
|
| 105 |
-
def generate_stream(self, domain, length=180, temperature=
|
| 106 |
-
|
|
|
|
|
|
|
|
|
|
| 107 |
m = self.models[domain]
|
| 108 |
ids = [self.bos] + (self.tok.encode(clean(prompt), add_special_tokens=False) if prompt else [])
|
| 109 |
t = torch.tensor([ids], device=self.dev)
|
| 110 |
-
bases = []
|
| 111 |
while sum(len(b) for b in bases) < length:
|
| 112 |
-
logits = m(input_ids=t).logits[:, -1, :] / max(temperature, 1e-6)
|
| 113 |
-
logits[:, :4] = -1e9
|
|
|
|
|
|
|
|
|
|
|
|
|
| 114 |
if top_k > 0:
|
| 115 |
v, _ = torch.topk(logits, top_k)
|
| 116 |
logits[logits < v[:, [-1]]] = -1e9
|
| 117 |
nxt = torch.multinomial(F.softmax(logits, dim=-1), 1)
|
|
|
|
| 118 |
t = torch.cat([t, nxt], dim=1)
|
| 119 |
-
bases.append(self.tok._ids_to_tokens[
|
| 120 |
yield "".join(bases)[:length]
|
| 121 |
|
| 122 |
-
def generate(self, domain, length=180, temperature=
|
|
|
|
| 123 |
out = ""
|
| 124 |
-
for out in self.generate_stream(domain, length, temperature, top_k, prompt):
|
| 125 |
pass
|
| 126 |
return out
|
|
|
|
| 102 |
return best, bpb
|
| 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=""):
|
| 107 |
+
"""Yield the growing continuation base-by-base (for live streaming UIs).
|
| 108 |
+
repetition_penalty discourages the homopolymer/repeat loops these small
|
| 109 |
+
specialists fall into; it does not erase genuine low-complexity domains."""
|
| 110 |
m = self.models[domain]
|
| 111 |
ids = [self.bos] + (self.tok.encode(clean(prompt), add_special_tokens=False) if prompt else [])
|
| 112 |
t = torch.tensor([ids], device=self.dev)
|
| 113 |
+
bases, emitted = [], []
|
| 114 |
while sum(len(b) for b in bases) < length:
|
| 115 |
+
logits = m(input_ids=t).logits[:, -1, :].float() / max(temperature, 1e-6)
|
| 116 |
+
logits[:, :4] = -1e9 # never emit specials
|
| 117 |
+
if repetition_penalty and repetition_penalty != 1.0: # penalize recent tokens
|
| 118 |
+
for tid in set(emitted[-12:]):
|
| 119 |
+
v = logits[0, tid]
|
| 120 |
+
logits[0, tid] = v / repetition_penalty if v > 0 else v * repetition_penalty
|
| 121 |
if top_k > 0:
|
| 122 |
v, _ = torch.topk(logits, top_k)
|
| 123 |
logits[logits < v[:, [-1]]] = -1e9
|
| 124 |
nxt = torch.multinomial(F.softmax(logits, dim=-1), 1)
|
| 125 |
+
ti = int(nxt); emitted.append(ti)
|
| 126 |
t = torch.cat([t, nxt], dim=1)
|
| 127 |
+
bases.append(self.tok._ids_to_tokens[ti])
|
| 128 |
yield "".join(bases)[:length]
|
| 129 |
|
| 130 |
+
def generate(self, domain, length=180, temperature=1.0, top_k=40,
|
| 131 |
+
repetition_penalty=1.3, prompt=""):
|
| 132 |
out = ""
|
| 133 |
+
for out in self.generate_stream(domain, length, temperature, top_k, repetition_penalty, prompt):
|
| 134 |
pass
|
| 135 |
return out
|