# 纯数据归因: 残差(over=参数预测名次-实际名次)在哪个层面聚集?
#   分组变量全部来自 CSV 现有列: 乐园 / 国家 / 机型 / 品牌前缀(机械提取)
#   指标: 组间方差占比 R2 (仅对 >=2 台在榜车的组), 置换检验 p 值
import csv, math, random, unicodedata, difflib

random.seed(7)
BASE = '.'
src_code = open(f'{BASE}/class05-bt-rank.py').read()
i0 = src_code.index('STEEL = [')
i1 = src_code.index("(50,'Twister','Knoebels')]") + len("(50,'Twister','Knoebels')]")
exec(src_code[i0:i1])

def parse_dur(s):
    if not s or ':' not in s: return None
    try:
        m, sec = s.split(':'); return int(m)*60 + int(sec)
    except ValueError: return None

def norm(s):
    s = unicodedata.normalize('NFKD', s)
    s = ''.join(c for c in s if not unicodedata.combining(c))
    return ''.join(c for c in s.lower() if c.isalnum())

rows = []
with open(f'{BASE}/rcdb_strict.csv') as f:
    for r in csv.DictReader(f):
        def num(k):
            try: return float(r[k])
            except (ValueError, TypeError): return None
        rows.append({'name': r['CoasterName'].split(' / ')[0], 'park': r['Park'], 'city': r['City'],
                     'country': r['Country'], 'status': r['Status'], 'type': r['Type'],
                     'design': r['Design'], 'year': r.get('OpSince', '')[:4],
                     'h': num('Height'), 's': num('Speed'), 'len': num('Length'),
                     'inv': num('Inversions'), 'dur': parse_dur(r.get('Duration', ''))})
for r in rows:
    r['nn'] = norm(r['name']); r['np'] = norm(r['park'] + r['city'])

def find(name, park):
    n, p = norm(name), norm(park)
    exact = [r for r in rows if r['nn'] == n]
    if len(exact) == 1: return exact[0]
    if len(exact) > 1:
        pk = [r for r in exact if p[:8] in r['np'] or r['np'][:8] in p]
        return pk[0] if pk else None
    sub = [r for r in rows if n in r['nn'] or (r['nn'] in n and len(r['nn']) >= 5)]
    if sub:
        pk = [r for r in sub if p[:6] in r['np']]
        if len(pk) == 1: return pk[0]
        if len(sub) == 1: return sub[0]
    best, score = None, 0
    for r in rows:
        sc = difflib.SequenceMatcher(None, n, r['nn']).ratio()
        if sc > score: best, score = r, sc
    if score >= 0.85 and (norm(park)[:5] in best['np'] or score >= 0.93):
        return best
    return None

usable = [r for r in rows if r['h'] and r['s'] and r['len']]
bothd = [(r['len']/r['s'], r['dur']) for r in usable if r['dur']]
n_ = len(bothd); mx = sum(x for x, _ in bothd)/n_; my = sum(y for _, y in bothd)/n_
tb = sum((x-mx)*(y-my) for x, y in bothd)/sum((x-mx)**2 for x, _ in bothd)
ta = my - tb*mx
for r in usable:
    r['dur_i'] = r['dur'] if r['dur'] else max(20.0, ta + tb*r['len']/r['s'])
    r['inv_i'] = r['inv'] if r['inv'] is not None else 0.0

FEATS = ['s', 'h', 'len', 'dur_i', 'inv_i']
stats = {}
for f_ in FEATS:
    v = [r[f_] for r in usable]
    m = sum(v)/len(v); sd = math.sqrt(sum((x-m)**2 for x in v)/len(v))
    stats[f_] = (m, sd)
def x_of(r): return [(r[f_]-stats[f_][0])/stats[f_][1] for f_ in FEATS]

def build(list_):
    out = []
    for rank, name, park in list_:
        r = find(name, park)
        if r and r['h'] and r['s'] and r['len']:
            out.append({'rank': rank, 'row': r})
    return out
steel = build(STEEL); wood = build(WOOD)

def pairs_of(lst):
    ps = []
    for i in range(len(lst)):
        for j in range(len(lst)):
            if lst[i]['rank'] < lst[j]['rank']:
                ps.append((x_of(lst[i]['row']), x_of(lst[j]['row'])))
    return ps

def fit(pairs, lam=0.02, iters=900, lr=0.1):
    k = len(FEATS); w = [0.0]*k
    for _ in range(iters):
        g = [0.0]*k
        for xa, xb in pairs:
            d = sum(w[t]*(xa[t]-xb[t]) for t in range(k))
            pr = 1/(1+math.exp(-max(-30, min(30, d))))
            for t in range(k):
                g[t] += (1-pr)*(xa[t]-xb[t])
        for t in range(k):
            w[t] += lr*(g[t]/len(pairs) - lam*w[t])
    return w

w = fit(pairs_of(steel) + pairs_of(wood))
for lst in (steel, wood):
    scored = sorted(lst, key=lambda e: -sum(w[t]*x_of(e['row'])[t] for t in range(len(w))))
    pr = {id(e['row']): i+1 for i, e in enumerate(scored)}
    for e in lst:
        e['over'] = pr[id(e['row'])] - e['rank']
allc = steel + wood
print(f'在榜车 {len(allc)} 台, over 均值 {sum(e["over"] for e in allc)/len(allc):+.1f}')

# ---- 分组方差占比 + 置换检验 ----
def brand_of(park):
    for b in ('Six Flags', 'Busch Gardens', 'Cedar Point', 'Kings ', 'Walibi', 'Knoebels', 'Kennywood'):
        if park.startswith(b): return b.strip()
    return park

def decomp(keyfn, label, min_n=2, B=3000):
    groups = {}
    for e in allc:
        groups.setdefault(keyfn(e['row']), []).append(e['over'])
    multi = {k: v for k, v in groups.items() if len(v) >= min_n}
    vals = [x for v in multi.values() for x in v]
    if len(multi) < 2 or len(vals) < 6:
        print(f'{label}: 组太少, 跳过'); return
    gm = sum(vals)/len(vals)
    sst = sum((x-gm)**2 for x in vals)
    def r2_of(assign):
        ssb = 0.0
        for v in assign:
            m = sum(v)/len(v)
            ssb += len(v)*(m-gm)**2
        return ssb/sst
    r2 = r2_of(list(multi.values()))
    sizes = [len(v) for v in multi.values()]
    cnt = 0
    pool = vals[:]
    for _ in range(B):
        random.shuffle(pool)
        idx = 0; assign = []
        for s_ in sizes:
            assign.append(pool[idx:idx+s_]); idx += s_
        if r2_of(assign) >= r2: cnt += 1
    p = (cnt+1)/(B+1)
    print(f'{label}: 组数={len(multi)}, 覆盖 {len(vals)} 台, 组间方差占比 R²={r2:.2f}, 置换检验 p={p:.3f}')
    top = sorted(multi.items(), key=lambda kv: -abs(sum(kv[1])/len(kv[1])))
    for k, v in top[:8]:
        print(f'    {str(k)[:30]:30s} n={len(v)}  组均值 {sum(v)/len(v):+6.1f}')

decomp(lambda r: r['park'], '按乐园 (>=2台在榜)')
decomp(lambda r: brand_of(r['park']), '按品牌前缀 (机械提取)')
decomp(lambda r: r['country'], '按国家')
decomp(lambda r: r['design'], '按机型')

# Six Flags 单独: 名称前缀纯机械判定
sf = [e['over'] for e in allc if e['row']['park'].startswith('Six Flags')]
nsf = [e['over'] for e in allc if not e['row']['park'].startswith('Six Flags')]
msf, mnsf = sum(sf)/len(sf), sum(nsf)/len(nsf)
diff = msf - mnsf
pool = sf + nsf; cnt = 0
for _ in range(5000):
    random.shuffle(pool)
    a = pool[:len(sf)]; b = pool[len(sf):]
    if abs(sum(a)/len(a) - sum(b)/len(b)) >= abs(diff): cnt += 1
print(f'\nSix Flags 前缀 (n={len(sf)}) over 均值 {msf:+.1f} vs 其他 (n={len(nsf)}) {mnsf:+.1f}, 差 {diff:+.1f}, 置换 p={(cnt+1)/5001:.3f}')
