File size: 2,109 Bytes
2ff4308
 
96bed25
 
 
 
 
 
 
2ff4308
 
 
 
 
 
 
 
 
 
96bed25
2ff4308
 
 
 
 
 
 
 
 
 
 
 
 
96bed25
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
import torch

from stego_olmoe.codebook import (
    build_frequency_balanced_code_values,
    codebook_bit_mass,
    decode_code_bits_to_token_ids,
    get_code_bits,
    make_codebook,
)


def test_codebook_deterministic():
    input_ids = torch.tensor([[1, 2, 3], [3, 2, 1]])
    bits_a = get_code_bits(input_ids, n_layers=16, vocab_size=10, seed=123)
    bits_b = get_code_bits(input_ids, n_layers=16, vocab_size=10, seed=123)
    assert torch.equal(bits_a, bits_b)


def test_codebook_balanced():
    codebook = make_codebook(vocab_size=128, n_layers=17, seed=7, code_scheme="balanced")
    ones = codebook.sum(dim=-1)
    assert torch.all(ones == 9)
    zeros = codebook.shape[-1] - ones
    assert torch.all((ones - zeros).abs() <= 1)


def test_codebook_shape_and_seed_changes():
    input_ids = torch.tensor([[0, 5, 11]])
    bits_a = get_code_bits(input_ids, n_layers=8, vocab_size=20, seed=1)
    bits_b = get_code_bits(input_ids, n_layers=8, vocab_size=20, seed=2)
    assert bits_a.shape == (1, 3, 8)
    assert not torch.equal(bits_a, bits_b)


def test_permuted_id_codebook_exactly_decodes_tokens():
    input_ids = torch.tensor([[0, 5, 11, 19]])
    bits = get_code_bits(input_ids, n_layers=8, vocab_size=20, seed=123, code_scheme="permuted_id")
    decoded = decode_code_bits_to_token_ids(bits, vocab_size=20, seed=123, code_scheme="permuted_id")
    assert torch.equal(decoded, input_ids)


def test_frequency_balanced_codebook_unique_balances_mass_and_decodes():
    counts = torch.tensor([100.0, 98.0, 30.0, 29.0, 5.0, 5.0, 1.0, 1.0])
    values = build_frequency_balanced_code_values(counts, n_layers=4, seed=0)
    assert torch.unique(values).numel() == counts.numel()

    bit_mass = codebook_bit_mass(values, counts, n_layers=4)
    assert torch.all((bit_mass - 0.5).abs() < 0.02)

    input_ids = torch.arange(counts.numel()).view(2, 4)
    bits = get_code_bits(input_ids, n_layers=4, vocab_size=counts.numel(), token_code_values=values)
    decoded = decode_code_bits_to_token_ids(bits, vocab_size=counts.numel(), token_code_values=values)
    assert torch.equal(decoded, input_ids)