Created
August 19, 2026 06:37
-
-
Save nishtahir/ed376a45b1673ae08a562518248db0e6 to your computer and use it in GitHub Desktop.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| import hashlib | |
| from itertools import pairwise | |
| import matplotlib.pyplot as plt | |
| import torch | |
| from transformers import AutoTokenizer, LogitsProcessor, Qwen3_5ForCausalLM | |
| GREEN = "\033[92m" | |
| RED = "\033[91m" | |
| RESET = "\033[0m" | |
| def green_ids( | |
| secret_key: str, prev_token_id: int, vocab_size: int, gamma: float | |
| ) -> torch.LongTensor: | |
| """Deterministically pick this step's green-list token ids from the secret key | |
| and previous token. The ordering of the returned ids is itself deterministic, | |
| which lets us split it into stable sub-partitions for message bits.""" | |
| seed_material = f"{secret_key}-{prev_token_id}".encode() | |
| digest = hashlib.sha256(seed_material).digest() | |
| # torch seeds need to fit in 64 bits — take the first 8 bytes | |
| seed = int.from_bytes(digest[:8], "big") | |
| generator = torch.Generator(device="cpu") | |
| generator.manual_seed(seed) | |
| perm = torch.randperm(vocab_size, generator=generator) | |
| cutoff = int(vocab_size * gamma) | |
| return perm[:cutoff] | |
| def text_to_bits(text: str) -> list[int]: | |
| data = text.encode("utf-8") | |
| return [(byte >> shift) & 1 for byte in data for shift in range(7, -1, -1)] | |
| def bits_to_text(bits: list[int]) -> str: | |
| n_bytes = len(bits) // 8 | |
| data = bytearray(n_bytes) | |
| for byte_index in range(n_bytes): | |
| value = 0 | |
| for bit in bits[byte_index * 8 : byte_index * 8 + 8]: | |
| value = (value << 1) | bit | |
| data[byte_index] = value | |
| try: | |
| # Prefer strict UTF-8 so correctly-decoded multi-byte text renders properly. | |
| return data.decode("utf-8") | |
| except UnicodeDecodeError: | |
| # A single flipped high bit turns an ASCII byte into a UTF-8 continuation | |
| # byte, which desyncs every byte after it into replacement characters. | |
| # Latin-1 is a 1:1 byte<->codepoint mapping, so on corrupted input each | |
| # bad bit shows up as one wrong character instead of cascading garbage. | |
| return data.decode("latin-1") | |
| class WatermarkWithEncodedMessage(LogitsProcessor): | |
| def __init__( | |
| self, | |
| secret_key: str, | |
| gamma: float = 0.5, | |
| delta: float = 1.0, | |
| message: str | None = None, | |
| message_delta: float = 1.0, | |
| ): | |
| self.secret_key = secret_key | |
| self.gamma = gamma | |
| self.delta = delta | |
| # Repeats across the whole generation (bit_index = step % len(message_bits)), | |
| # so the decoder can recover each bit by majority vote over many observations. | |
| self.message_bits = text_to_bits(message) if message else None | |
| self.message_delta = message_delta | |
| self.step = 0 | |
| def __call__( | |
| self, | |
| input_ids: torch.LongTensor, | |
| scores: torch.FloatTensor, | |
| ) -> torch.FloatTensor: | |
| batch_size, vocab_size = scores.shape | |
| target_bit = None | |
| if self.message_bits: | |
| target_bit = self.message_bits[self.step % len(self.message_bits)] | |
| for i in range(batch_size): | |
| prev_token_id = input_ids[i, -1].item() | |
| ids = green_ids(self.secret_key, prev_token_id, vocab_size, self.gamma) | |
| # Boost the probability of green token ids | |
| scores[i, ids] += self.delta | |
| if target_bit is not None: | |
| # Additionally boost whichever green half encodes the current | |
| # message bit, without touching the overall green/red split. | |
| half = len(ids) // 2 | |
| signal_ids = ids[half:] if target_bit else ids[:half] | |
| scores[i, signal_ids] += self.message_delta | |
| self.step += 1 | |
| return scores | |
| def detect( | |
| token_ids: list[int], secret_key: str, vocab_size: int, gamma: float | |
| ) -> tuple[list[bool], list[int | None]]: | |
| """For each token after the first, report whether it fell in that step's green | |
| list, and if so which half (0/1) it landed in — a vote for that step's message | |
| bit, or None if the token was red.""" | |
| green_flags = [] | |
| bit_votes = [] | |
| for prev_id, token_id in pairwise(token_ids): | |
| ids = green_ids(secret_key, prev_id, vocab_size, gamma) | |
| half = len(ids) // 2 | |
| if token_id in ids[:half]: | |
| green_flags.append(True) | |
| bit_votes.append(0) | |
| elif token_id in ids[half:]: | |
| green_flags.append(True) | |
| bit_votes.append(1) | |
| else: | |
| green_flags.append(False) | |
| bit_votes.append(None) | |
| return green_flags, bit_votes | |
| def decode_message_bits(votes: list[int | None], num_bits: int) -> list[int]: | |
| """Majority-vote each message bit from its (possibly many) repeated observations.""" | |
| decoded = [] | |
| for bit_index in range(num_bits): | |
| observations = [ | |
| vote | |
| for step, vote in enumerate(votes) | |
| if vote is not None and step % num_bits == bit_index | |
| ] | |
| ones = sum(observations) | |
| zeros = len(observations) - ones | |
| decoded.append(1 if ones >= zeros else 0) | |
| return decoded | |
| def z_scores(flags: list[bool], gamma: float) -> list[float]: | |
| """Cumulative watermark detection z-score after each observed token.""" | |
| scores = [] | |
| green_count = 0 | |
| for n, flag in enumerate(flags, start=1): | |
| green_count += flag | |
| expected = gamma * n | |
| std = (n * gamma * (1 - gamma)) ** 0.5 | |
| scores.append((green_count - expected) / std if std > 0 else 0.0) | |
| return scores | |
| def summarize_flags(flags: list[bool]) -> str: | |
| green_count = sum(flags) | |
| red_count = len(flags) - green_count | |
| green_pct = 100 * green_count / len(flags) if flags else 0.0 | |
| return f"{green_count} green / {red_count} red ({green_pct:.1f}% green)" | |
| def render_highlighted(tokenizer, token_ids: list[int], flags: list[bool]) -> str: | |
| """First token plain, then each following token colored green/red by watermark membership.""" | |
| pieces = [ | |
| tokenizer.convert_tokens_to_string( | |
| [tokenizer.convert_ids_to_tokens(token_ids[0])] | |
| ) | |
| ] | |
| for token_id, flag in zip(token_ids[1:], flags): | |
| text = tokenizer.convert_tokens_to_string( | |
| [tokenizer.convert_ids_to_tokens(token_id)] | |
| ) | |
| color = GREEN if flag else RED | |
| pieces.append(f"{color}{text}{RESET}") | |
| return "".join(pieces) | |
| def plot_z_scores(watermarked: list[float], plain: list[float], path: str): | |
| fig, ax = plt.subplots(figsize=(8, 5)) | |
| ax.plot( | |
| range(1, len(watermarked) + 1), | |
| watermarked, | |
| label="Watermarked", | |
| color="tab:green", | |
| ) | |
| ax.plot(range(1, len(plain) + 1), plain, label="Plain", color="tab:gray") | |
| ax.axhline( | |
| 4.0, | |
| color="tab:red", | |
| linestyle="--", | |
| linewidth=1, | |
| label="z = 4 (typical detection threshold)", | |
| ) | |
| ax.set_xlabel("Generated token index") | |
| ax.set_ylabel("Cumulative detection z-score") | |
| ax.set_title("Watermark strength over the generated sequence") | |
| ax.legend() | |
| fig.tight_layout() | |
| fig.savefig(path, dpi=150) | |
| print(f"Saved z-score plot to {path}") | |
| def report( | |
| label: str, | |
| tokenizer, | |
| output_ids: torch.LongTensor, | |
| prompt_len: int, | |
| secret_key: str, | |
| vocab_size: int, | |
| gamma: float, | |
| message_bits: list[int], | |
| ) -> list[float]: | |
| """Print the generated text, watermark strength, and decoded message for one | |
| generation, returning its cumulative z-score series for plotting.""" | |
| # Include the last prompt token so the first generated token's green list can be recomputed | |
| token_ids = output_ids[0, prompt_len - 1 :].tolist() | |
| flags, votes = detect(token_ids, secret_key, vocab_size, gamma) | |
| z = z_scores(flags, gamma) | |
| decoded = decode_message_bits(votes, len(message_bits)) | |
| accuracy = sum(a == b for a, b in zip(message_bits, decoded)) / len(message_bits) | |
| print(f"\n{label}:") | |
| print(tokenizer.decode(output_ids[0, prompt_len:], skip_special_tokens=True)) | |
| print(render_highlighted(tokenizer, token_ids, flags)) | |
| print(summarize_flags(flags)) | |
| print(f"Final z-score: {z[-1]:.2f}") | |
| print(f"Decoded message: {bits_to_text(decoded)!r} (bit accuracy: {accuracy:.1%})") | |
| return z | |
| def main(): | |
| device = torch.device("mps" if torch.backends.mps.is_available() else "cpu") | |
| model_name = "Qwen/Qwen3.5-4B" | |
| secret_key = "leftrightba" | |
| gamma = 0.5 | |
| secret_message = "hello world" | |
| message_delta = 4.0 | |
| max_new_tokens = 600 | |
| temperature = 0.1 | |
| model = Qwen3_5ForCausalLM.from_pretrained(model_name, device_map={"": device}) | |
| tokenizer = AutoTokenizer.from_pretrained(model_name) | |
| vocab_size = model.config.vocab_size | |
| messages = [ | |
| { | |
| "role": "user", | |
| "content": "Summarize the plot of the hounds of baskerville.", | |
| } | |
| ] | |
| inputs = tokenizer.apply_chat_template( | |
| messages, | |
| tokenize=True, | |
| enable_thinking=False, # Disables the <think> block | |
| add_generation_prompt=True, | |
| return_dict=True, | |
| return_tensors="pt", | |
| ) | |
| inputs = inputs.to(device) | |
| prompt_len = inputs["input_ids"].shape[1] | |
| message_bits = text_to_bits(secret_message) | |
| # Only green tokens (~gamma of steps) cast a vote, and the message cycles once | |
| # per len(message_bits) steps — this estimates how many votes each bit gets. | |
| expected_votes_per_bit = (max_new_tokens / len(message_bits)) * gamma | |
| if expected_votes_per_bit < 8: | |
| print( | |
| f"Warning: {secret_message!r} is {len(message_bits)} bits but " | |
| f"max_new_tokens={max_new_tokens} gives only ~{expected_votes_per_bit:.1f} " | |
| "green-token observations per bit on average — too little redundancy for " | |
| "majority vote to correct errors reliably. Shorten the message or raise " | |
| "max_new_tokens (aim for at least ~8 observations per bit)." | |
| ) | |
| outputs_watermarked = model.generate( | |
| **inputs, | |
| max_new_tokens=max_new_tokens, | |
| do_sample=True, | |
| temperature=temperature, | |
| logits_processor=[ | |
| Watermarking( | |
| secret_key=secret_key, | |
| gamma=gamma, | |
| message=secret_message, | |
| message_delta=message_delta, | |
| ), | |
| ], | |
| ) | |
| outputs_plain = model.generate( | |
| **inputs, | |
| max_new_tokens=max_new_tokens, | |
| do_sample=True, | |
| temperature=temperature, | |
| ) | |
| watermarked_z = report( | |
| "Watermarked", | |
| tokenizer, | |
| outputs_watermarked, | |
| prompt_len, | |
| secret_key, | |
| vocab_size, | |
| gamma, | |
| message_bits, | |
| ) | |
| plain_z = report( | |
| "Plain", | |
| tokenizer, | |
| outputs_plain, | |
| prompt_len, | |
| secret_key, | |
| vocab_size, | |
| gamma, | |
| message_bits, | |
| ) | |
| plot_z_scores(watermarked_z, plain_z, "watermark_analysis.png") | |
| if __name__ == "__main__": | |
| main() |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment