Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 20 additions & 2 deletions perceval/utils/conversion.py
Original file line number Diff line number Diff line change
Expand Up @@ -174,11 +174,29 @@ def sample_count_to_samples(sample_count: BSCount, **kwargs) -> BSSamples:

:return: the sample list
"""
total = sample_count.total()
try:
count = _deduce_count(**kwargs)
except RuntimeError:
count = sum(sample_count.values())
return sample_count_to_probs(sample_count).sample(count, non_null=False)
count = total
if count < 0:
raise RuntimeError(f"A sample count must be positive (got {count})")
# Shuffle then keep only the requested number of samples
sample_list = [
state
for state, nb in sample_count.items()
for _ in range(nb)
]
random.shuffle(sample_list)
samples = BSSamples()
samples.extend(sample_list[:count])
# If more samples requested then do random sampling
if count > total:
samples.extend(
sample_count_to_probs(sample_count).sample(count - total,
non_null=False)
)
return samples


class ConversionHelper:
Expand Down
16 changes: 12 additions & 4 deletions tests/utils/test_utils_conversion.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,16 +96,24 @@ def test_probs_to_sample_count(count):
assert sum(list(output.values())) == count


def test_sample_count_to_samples():
@pytest.mark.parametrize("count", [1000, 1000000, 170, 1, 0])
def test_sample_count_to_samples(count):
sample_count = BSCount({
b0: 280,
b1: 120,
b2: 400,
b3: 200
})
samples = sample_count_to_samples(sample_count)
for state, count in sample_count.items():
assert count * 0.7 < samples.count(state) < count * 1.3
samples = sample_count_to_samples(sample_count, max_samples=int(count))
if count == 1000:
assert samples.count(b0) == 280
assert samples.count(b1) == 120
assert samples.count(b2) == 400
assert samples.count(b3) == 200
if count > 1000:
for _state, _count in samples_to_sample_count(samples).items():
assert _count * 0.7 < samples.count(_state) < _count * 1.3
assert len(samples) == count


def test_probs_to_samples():
Expand Down
Loading