Skip to content

联合探查

在预算内比较控制组合、缩尾和样本处理。

核心代码

py
def build_followup(adapter, c, state, records, remaining):
    """结果更新在批次边界进行,恢复不会因完成先后产生另一条搜索路径。"""
    if remaining <= 0 or c['mode'] == 'exhaustive':
        return []
    valid = sorted([r for r in records if r['result'].get('status') == 'ok'], key=lambda r: objective(c, r))
    seen = {r['spec']['key'] for r in records}
    proposals = []
    if state['stage'] == 'initial' and 'joint' in c['strategies']:
        top = valid[:10]
        control_seeds = list(dict.fromkeys(tuple(r['spec']['controls']) for r in top)) or [tuple(c['baseline_controls'])]
        if 'controls' in c['strategies']:
            control_seeds += [tuple(v) for v in control_options(c, 10)]
        levels = c['winsor_levels'] if 'winsor' in c['strategies'] else [0.]
        samples = [r['spec']['sample'] for r in top if r['spec']['sample']['kind'] != 'none']
        extra, _complete = sample_options(adapter, c, max(3, remaining//20))
        samples += [v for _strategy, v in extra[:30]]
        samples = [{'kind': 'none'}] + samples
        combos = list(itertools.product(control_seeds, levels, samples))
        rng = np.random.default_rng(c['seed']+1)
        rng.shuffle(combos)
        proposals += [candidate(list(a), b, d, 'joint', True) for a,b,d in combos[:max(1, remaining*2//3)]]
        # 联合随机探索不依赖单策略得分,保留发现交互改善的机会。
        for i in range(max(1, remaining//5)):
            controls = [v for g in c['groups'] if rng.random() < .5 for v in g] if 'controls' in c['strategies'] else c['baseline_controls']
            if controls_valid(c, controls):
                proposals.append(candidate(controls, float(rng.choice(levels)), samples[int(rng.integers(len(samples)))], 'joint', True))
    else:
        for record in valid[:10]:
            s = record['spec']
            if 'controls' in c['strategies'] and ('joint' in c['strategies'] or (s['winsor']==0 and s['sample']['kind']=='none')):
                chosen = set(s['controls'])
                for group in c['groups']:
                    changed = sorted(chosen ^ set(group))
                    if controls_valid(c, changed):
                        proposals.append(candidate(changed,s['winsor'],s['sample'],'joint' if 'joint' in c['strategies'] else 'controls',True))
                for remove, add in itertools.product(c['groups'], repeat=2):
                    changed = sorted((chosen-set(remove)) | set(add))
                    if controls_valid(c, changed):
                        proposals.append(candidate(changed,s['winsor'],s['sample'],'joint' if 'joint' in c['strategies'] else 'controls',True))
            if s['sample']['kind'] in ('drop_rows', 'drop_entities'):
                values = s['sample']['values']
                pool = adapter.ref_ids if s['sample']['kind']=='drop_rows' else sorted(adapter.reference[c['entity']].astype(str).unique())
                for old in values:
                    proposals.append(candidate(s['controls'],s['winsor'],{'kind':s['sample']['kind'],'values':[v for v in values if v!=old]},s['strategy'],True))
                if len(values) < c['delete_depth']:
                    rng = np.random.default_rng(c['seed']+len(values))
                    for index in rng.choice(len(pool), min(len(pool), 30), replace=False):
                        value = pool[int(index)]
                        if value not in values:
                            proposals.append(candidate(s['controls'],s['winsor'],{'kind':s['sample']['kind'],'values':sorted(values+[value])},s['strategy'],True))
            if 'winsor' in c['strategies'] and ('joint' in c['strategies'] or (s['controls']==sorted(c['baseline_controls']) and s['sample']['kind']=='none')):
                for level in c['winsor_levels']:
                    proposals.append(candidate(s['controls'],level,s['sample'],'joint' if 'joint' in c['strategies'] else 'winsor',True))
    return [s for s in merge_specs(proposals) if s['key'] not in seen][:remaining]

Released under the AGPL-3.0 License.