Spaces:
Running on Zero
Running on Zero
Download byo_features.py from jaydeepb/predicting_memorization: direct link, hf CLI and curl.
- Browser
- Download file 3.89 kB
-
https://huggingface.co/spaces/jaydeepb/predicting_memorization/resolve/main/byo_features.py
- Command line
-
hf download hf://spaces/jaydeepb/predicting_memorization/byo_features.py
-
curl -L -o byo_features.py https://huggingface.co/spaces/jaydeepb/predicting_memorization/resolve/main/byo_features.py
3.89 kB
| """Compute the six base-model features for an arbitrary user sequence, mirroring | |
| feature_extraction_{cpu,gpu}.py so scores are comparable to the paper's classifier. | |
| Definitions (base model, no fine-tuning): | |
| zlib_entropy = compressed byte length of the decoded 50-token suffix | |
| base_ppl = exp(mean per-token loss) over the 50-token continuation | |
| base_loss_variance = variance of that per-token loss | |
| gradient_norm = ||grad of the continuation loss w.r.t. all params||_2 | |
| one_step_abs_change = loss_before - loss_after, after one SGD step (lr=2e-5) on x | |
| one_step_rel_change = one_step_abs_change / loss_before | |
| """ | |
| import zlib, numpy as np, torch, torch.nn.functional as F | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| PREFIX_LEN, WINDOW_LEN, LR = 50, 100, 2e-5 | |
| class BYOExtractor: | |
| def __init__(self, model_path, device=None): | |
| self.device = device or ("cuda" if torch.cuda.is_available() else "cpu") | |
| self.tok = AutoTokenizer.from_pretrained(model_path) | |
| self.model = AutoModelForCausalLM.from_pretrained( | |
| model_path, dtype=torch.float32).to(self.device).eval() | |
| def _cont_loss(self, ids): | |
| logits = self.model(ids).logits | |
| cl = logits[:, PREFIX_LEN - 1:WINDOW_LEN - 1, :] | |
| ct = ids[:, PREFIX_LEN:WINDOW_LEN] | |
| return -F.log_softmax(cl, -1).gather(2, ct.unsqueeze(-1)).squeeze(-1) # [1,50] per-token loss | |
| def features(self, text): | |
| enc = self.tok(text, return_tensors="pt", truncation=True, max_length=WINDOW_LEN) | |
| ids = enc["input_ids"] | |
| if ids.shape[1] < WINDOW_LEN: | |
| raise ValueError( | |
| f"The input should be at least {WINDOW_LEN} tokens (or about 75 words); " | |
| f"you entered {ids.shape[1]} tokens.") | |
| return self.features_from_ids(ids[:, :WINDOW_LEN]) | |
| def features_from_ids(self, ids): | |
| """ids: LongTensor [1, WINDOW_LEN] of exact token ids (bypasses tokenization). | |
| Used both by features(text) and by the offline validation against the pipeline.""" | |
| if not torch.is_tensor(ids): | |
| ids = torch.tensor(ids, dtype=torch.long) | |
| if ids.dim() == 1: | |
| ids = ids.unsqueeze(0) | |
| ids = ids[:, :WINDOW_LEN].long().to(self.device) | |
| # zlib on the decoded suffix | |
| suffix_txt = self.tok.decode(ids[0, PREFIX_LEN:WINDOW_LEN], skip_special_tokens=True) | |
| zlib_entropy = float(len(zlib.compress(suffix_txt.encode("utf-8")))) | |
| # loss-based features | |
| with torch.no_grad(): | |
| ptl = self._cont_loss(ids) | |
| base_ppl = float(torch.exp(ptl.mean())) | |
| base_loss_variance = float(ptl.var(unbiased=True)) | |
| # gradient norm + one-step loss change | |
| self.model.zero_grad(set_to_none=True) | |
| loss_before = self._cont_loss(ids).mean() | |
| loss_before.backward() | |
| with torch.no_grad(): | |
| gsq = sum(float(p.grad.pow(2).sum()) for p in self.model.parameters() if p.grad is not None) | |
| gradient_norm = gsq ** 0.5 | |
| for p in self.model.parameters(): # one SGD step | |
| if p.grad is not None: | |
| p.data.add_(p.grad, alpha=-LR) | |
| loss_after = float(self._cont_loss(ids).mean()) | |
| for p in self.model.parameters(): # restore | |
| if p.grad is not None: | |
| p.data.add_(p.grad, alpha=LR) | |
| lb = float(loss_before) | |
| one_step_abs_change = lb - loss_after | |
| one_step_rel_change = one_step_abs_change / lb if lb else 0.0 | |
| self.model.zero_grad(set_to_none=True) | |
| return { | |
| "zlib_entropy": zlib_entropy, "base_ppl": base_ppl, | |
| "base_loss_variance": base_loss_variance, "gradient_norm": gradient_norm, | |
| "one_step_abs_change": one_step_abs_change, "one_step_rel_change": one_step_rel_change, | |
| } | |