Skip to content

Instantly share code, notes, and snippets.

@nishtahir
Created August 19, 2026 06:37
Show Gist options
  • Select an option

  • Save nishtahir/ed376a45b1673ae08a562518248db0e6 to your computer and use it in GitHub Desktop.

Select an option

Save nishtahir/ed376a45b1673ae08a562518248db0e6 to your computer and use it in GitHub Desktop.
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