stego-olmoe-router-code / tests /test_alignment.py
anpaurehf's picture
Upload router-code SFT repo
96bed25 verified
Raw
History Blame Contribute Delete
765 Bytes
import torch
from stego_olmoe.codebook import get_code_bits
def test_router_targets_align_to_current_input_not_next_label():
input_ids = torch.tensor([[101, 202, 303]])
labels = torch.tensor([[202, 303, 404]])
router_targets = get_code_bits(input_ids, n_layers=10, vocab_size=1000, seed=5)
label_targets = get_code_bits(labels, n_layers=10, vocab_size=1000, seed=5)
assert torch.equal(router_targets[0, 0], get_code_bits(torch.tensor([[101]]), 10, 1000, seed=5)[0, 0])
assert torch.equal(router_targets[0, 1], get_code_bits(torch.tensor([[202]]), 10, 1000, seed=5)[0, 0])
assert torch.equal(router_targets[0, 2], get_code_bits(torch.tensor([[303]]), 10, 1000, seed=5)[0, 0])
assert not torch.equal(router_targets, label_targets)