MCPcopy Create free account
hub / github.com/bamler-lab/constriction / test_chain_gaussian

Function test_chain_gaussian

tests/python/test_constriction.py:58–99  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

56
57
58def test_chain_gaussian():
59 rng = np.random.RandomState(123)
60 original_data = rng.randint(2**32, size=100, dtype=np.uint32)
61 decoder = constriction.stream.chain.ChainCoder(original_data, seal=True)
62
63 model = constriction.stream.model.QuantizedGaussian(-100, 100)
64 means = np.arange(50, dtype=np.float64)
65 stds = np.array([10.0] * 50, dtype=np.float64)
66
67 symbols = decoder.decode(model, means, stds)
68
69 remainders_prefix, remainders_suffix = decoder.get_remainders()
70 print(len(remainders_prefix), len(remainders_suffix), len(original_data))
71 assert len(remainders_prefix) + len(remainders_suffix) < len(original_data)
72
73 # Variant 1: treat `remainders_prefix` and `remainders_suffix` separately
74 encoder1 = constriction.stream.chain.ChainCoder(
75 remainders_suffix, is_remainders=True)
76 encoder1.encode_reverse(symbols, model, means, stds)
77 recovered_prefix1, recovered_suffix1 = encoder1.get_data(unseal=True)
78 print(len(recovered_prefix1), len(recovered_suffix1), len(original_data))
79 assert len(recovered_prefix1) == 0
80 recovered1 = np.concatenate((remainders_prefix, recovered_suffix1))
81 assert np.all(recovered1 == original_data)
82
83 # Variant 2: concatenate `remainders_prefix` and `remainders_suffix`
84 remainders = np.concatenate((remainders_prefix, remainders_suffix))
85 encoder2 = constriction.stream.chain.ChainCoder(
86 remainders, is_remainders=True)
87 encoder2.encode_reverse(symbols, model, means, stds)
88 recovered_prefix2, recovered_suffix2 = encoder2.get_data(unseal=True)
89 print(len(recovered_prefix2), len(recovered_suffix2), len(original_data))
90 recovered2 = np.concatenate((recovered_prefix2, recovered_suffix2))
91 assert np.all(recovered2 == original_data)
92
93 # Variant 3: directly re-encode onto original coder
94 encoder3 = decoder
95 encoder3.encode_reverse(symbols, model, means, stds)
96 recovered_prefix3, recovered_suffix3 = encoder3.get_data(unseal=True)
97 print(len(recovered_prefix3), len(recovered_suffix3), len(original_data))
98 assert len(recovered_prefix3) == 0
99 assert np.all(recovered_suffix3 == original_data)
100
101
102def test_chain_independence():

Callers

nothing calls this directly

Calls 5

decodeMethod · 0.95
get_remaindersMethod · 0.95
encode_reverseMethod · 0.95
get_dataMethod · 0.95
encode_reverseMethod · 0.45

Tested by

no test coverage detected