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)