Quazim0t0 commited on
Commit
1a5e7a2
·
verified ·
1 Parent(s): 92b7758

Upload daisychain.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. 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=0.9, top_k=20, prompt=""):
106
- """Yield the growing continuation base-by-base (for live streaming UIs)."""
 
 
 
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[int(nxt)])
120
  yield "".join(bases)[:length]
121
 
122
- def generate(self, domain, length=180, temperature=0.9, top_k=20, prompt=""):
 
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