22 lines
728 B
Python
22 lines
728 B
Python
from __future__ import annotations
|
|
|
|
import random
|
|
|
|
|
|
def relative_high_low(
|
|
score_by_id: dict[str, float], count: int = 3, seed: str | int = 0
|
|
) -> tuple[list[str], list[str]]:
|
|
"""稳定选择互不重叠的相对高分组和低分组。"""
|
|
|
|
if len(score_by_id) < count * 2:
|
|
raise ValueError(f"need at least {count * 2} traces for disjoint High/Low groups")
|
|
trace_ids = list(score_by_id)
|
|
random.Random(str(seed)).shuffle(trace_ids)
|
|
high = sorted(trace_ids, key=score_by_id.__getitem__, reverse=True)[:count]
|
|
high_set = set(high)
|
|
low = sorted(
|
|
(trace_id for trace_id in trace_ids if trace_id not in high_set),
|
|
key=score_by_id.__getitem__,
|
|
)[:count]
|
|
return high, low
|