# 监督版 2018A: 用 Golden Ticket 2018 两张 top-50 当标签
# 1) 榜单条目匹配到 rcdb_strict 行 (名称+乐园模糊匹配, 打印全部匹配供人工核查)
# 2) 榜内成对比较 + pairwise logistic (Bradley-Terry), L2 正则, 纯 python 梯度下降
# 3) 5 折交叉验证 (按车分折) 报 Kendall tau; 对照单变量基线
# 4) bootstrap 系数区间; 5) 全库 1207 台预测分 -> 低估/高估名单
import csv, math, random, unicodedata, difflib

random.seed(7)
SRC = 'rcdb_strict.csv'

STEEL = [(1,'Fury 325','Carowinds'),(2,'Millennium Force','Cedar Point'),(3,'Steel Vengeance','Cedar Point'),
(4,'Expedition GeForce','Holiday Park'),(5,'Superman: The Ride','Six Flags New England'),
(6,"Apollo's Chariot",'Busch Gardens Williamsburg'),(7,'Iron Rattler','Six Flags Fiesta Texas'),
(8,'Leviathan',"Canada's Wonderland"),(9,'Maverick','Cedar Point'),(10,'Diamondback','Kings Island'),
(11,'Nitro','Six Flags Great Adventure'),(12,'Intimidator 305','Kings Dominion'),
(13,"Phantom's Revenge",'Kennywood'),(14,'Magnum XL-200','Cedar Point'),(15,'Taron','Phantasialand'),
(16,'Top Thrill Dragster','Cedar Point'),(17,'Mako','SeaWorld Orlando'),(18,'Time Traveler','Silver Dollar City'),
(19,'Blue Fire','Europa Park'),(20,'Nemesis','Alton Towers'),(21,'Helix','Liseberg'),
(22,'Intimidator','Carowinds'),(23,'New Texas Giant','Six Flags Over Texas'),
(24,'Twisted Colossus','Six Flags Magic Mountain'),(25,'Mind Bender','Six Flags Over Georgia'),
(26,'Goliath','Six Flags Over Georgia'),(27,'Behemoth',"Canada's Wonderland"),(28,'Montu','Busch Gardens Tampa'),
(29,'Banshee','Kings Island'),(30,'Skyrush','Hersheypark'),(31,'X2','Six Flags Magic Mountain'),
(32,'Alpengeist','Busch Gardens Williamsburg'),(33,'Wicked Cyclone','Six Flags New England'),
(34,'Black Mamba','Phantasialand'),(35,'Cheetah Hunt','Busch Gardens Tampa'),
(36,'Verbolten','Busch Gardens Williamsburg'),(37,'Kumba','Busch Gardens Tampa'),
(38,'Twisted Timbers','Kings Dominion'),(39,'Jetline','Grona Lund'),(39,'Superman Ride of Steel','Six Flags America'),
(41,'Goliath','La Ronde'),(42,'Lisebergbanan','Liseberg'),(43,'Griffon','Busch Gardens Williamsburg'),
(44,'Cannibal','Lagoon'),(45,'Shambhala','PortAventura'),(46,'Expedition Everest',"Disney's Animal Kingdom"),
(47,'Storm Chaser','Kentucky Kingdom'),(48,'Raging Bull','Six Flags Great America'),
(49,'Thunderbird','Holiday World'),(50,'Whizzer','Six Flags Great America')]

WOOD = [(1,'Phoenix','Knoebels'),(2,'El Toro','Six Flags Great Adventure'),(3,'Voyage','Holiday World'),
(4,'Boulder Dash','Lake Compounce'),(5,'Beast','Kings Island'),(6,'Lightning Rod','Dollywood'),
(7,'Outlaw Run','Silver Dollar City'),(8,'Ravine Flyer II','Waldameer'),(9,'Gold Striker',"California's Great America"),
(10,'Thunderhead','Dollywood'),(11,'Mystic Timbers','Kings Island'),(12,'Lightning Racer','Hersheypark'),
(13,'GhostRider',"Knott's Berry Farm"),(14,'Balder','Liseberg'),(15,'Thunderbolt','Kennywood'),
(16,'Wodan','Europa Park'),(17,'Wildfire','Kolmarden'),(18,'Raven','Holiday World'),
(19,'Goliath','Six Flags Great America'),(20,'Jack Rabbit','Kennywood'),
(21,'Giant Dipper','Santa Cruz Beach Boardwalk'),(21,'Shivering Timbers',"Michigan's Adventure"),
(23,'Legend','Holiday World'),(24,'White Lightning','Fun Spot'),(25,'Troy','Toverland'),
(26,'Renegade','Valleyfair'),(27,'Cu Chulainn','Tayto Park'),(28,'Colossos','Heide Park'),
(29,'Cyclone','Luna Park Coney Island'),(30,'Prowler','Worlds of Fun'),(30,'Rampage','Alabama Splash Adventure'),
(32,'Comet','Great Escape'),(33,'Boardwalk Bullet','Kemah Boardwalk'),(34,'Flying Turns','Knoebels'),
(35,'Rutschebanan','Tivoli Gardens'),(36,'Switchback','ZDTs'),(37,'Blue Streak','Conneaut Lake'),
(38,'American Thunder','Six Flags St. Louis'),(38,"Screamin' Eagle",'Six Flags St. Louis'),
(40,'Grizzly','Kings Dominion'),(41,'Blue Streak','Cedar Point'),(42,'Playland Wooden Coaster','Playland'),
(43,'Boss','Six Flags St. Louis'),(44,'Wild One','Six Flags America'),(45,'T Express','Everland'),
(46,'Megafobia','Oakwood'),(47,'Hades 360','Mount Olympus'),(48,'Mine Blower','Fun Spot'),
(49,'Wooden Warrior','Quassy'),(50,'Twister','Knoebels')]

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(SRC) 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'],
                     '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], 'exact'
    if len(exact) > 1:
        pk = [r for r in exact if p[:8] in r['np'] or r['np'][:8] in p or
              any(w in r['np'] for w in [p[:6]] if len(w) >= 5)]
        if len(pk) >= 1: return pk[0], 'exact+park'
        return None, 'ambiguous'
    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], 'substr+park'
        if len(sub) == 1: return sub[0], 'substr'
    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, f'fuzzy{score:.2f}'
    return None, f'no(best={best["name"]}@{best["park"]} {score:.2f})'

# 时长推补 (同前口径, 在全库上拟合)
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']
# 标准化参数取自全库 usable
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_, label):
    out, miss = [], []
    for rank, name, park in list_:
        r, how = find(name, park)
        if r is None or r not in usable and not (r and r['h'] and r['s'] and r['len']):
            miss.append((rank, name, park, how)); continue
        if not (r['h'] and r['s'] and r['len']):
            miss.append((rank, name, park, '字段缺')); continue
        if r['dur'] is None and 'dur_i' not in r:
            miss.append((rank, name, park, '时长不可推')); continue
        out.append({'rank': rank, 'gname': name, 'row': r, 'how': how})
    print(f'\n== {label}: 匹配 {len(out)}/{len(list_)} ==')
    for e in out:
        r = e['row']
        flag = '' if norm(e['gname']) == r['nn'] else f"  <- 匹配为 {r['name']}@{r['park']} ({e['how']})"
        if flag: print(f"  #{e['rank']:2d} {e['gname']}{flag}")
    for m in miss:
        print(f"  #{m[0]:2d} {m[1]} @ {m[2]}  未匹配: {m[3]}")
    return out

steel = build(STEEL, '钢质榜')
wood = build(WOOD, '木质榜')

# ---- BT 拟合 ----
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=1500, 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

def tau(lst, w):
    # Kendall tau: 预测序 vs 榜单序
    c = d = 0
    for i in range(len(lst)):
        for j in range(i+1, len(lst)):
            a, b = lst[i], lst[j]
            if a['rank'] == b['rank']: continue
            sa = sum(w[t]*x_of(a['row'])[t] for t in range(len(w)))
            sb = sum(w[t]*x_of(b['row'])[t] for t in range(len(w)))
            agree = (a['rank'] < b['rank']) == (sa > sb)
            c += agree; d += (not agree)
    return (c-d)/(c+d)

all_pairs = pairs_of(steel) + pairs_of(wood)
w = fit(all_pairs)
print(f'\n== BT 系数 (标准化) == 样本对 {len(all_pairs)}')
for f_, wi in zip(FEATS, w):
    print(f'  {f_:6s} {wi:+.3f}')
print(f'  极速+高度 组合 {w[0]+w[1]:+.3f}')
print(f'in-sample tau: 钢 {tau(steel, w):+.3f}  木 {tau(wood, w):+.3f}')

# 单变量基线
for bi, f_ in enumerate(FEATS):
    wb = [0.0]*len(FEATS); wb[bi] = 1.0
    print(f'  基线 只用{f_:6s}: 钢 tau {tau(steel, wb):+.3f}  木 {tau(wood, wb):+.3f}')

# ---- 5 折 CV ----
def cv(k=5):
    lab = [(e, 'S') for e in steel] + [(e, 'W') for e in wood]
    random.shuffle(lab)
    folds = [lab[i::k] for i in range(k)]
    taus = []
    for fd in folds:
        held = set(id(e['row']) for e, _ in fd)
        tr_s = [e for e in steel if id(e['row']) not in held]
        tr_w = [e for e in wood if id(e['row']) not in held]
        wf = fit(pairs_of(tr_s) + pairs_of(tr_w), iters=800)
        te_s = [e for e in steel if id(e['row']) in held]
        te_w = [e for e in wood if id(e['row']) in held]
        for te in (te_s, te_w):
            if len(te) >= 5:
                taus.append(tau(te, wf))
    print(f'5折CV held-out tau: 均值 {sum(taus)/len(taus):+.3f}  各折 {[f"{t:+.2f}" for t in taus]}')
cv()

# ---- bootstrap 系数 ----
B = 100
bw = []
for _ in range(B):
    bs = [random.choice(steel) for _ in steel]
    bwd = [random.choice(wood) for _ in wood]
    bw.append(fit(pairs_of(bs) + pairs_of(bwd), iters=600))
print('\n== bootstrap 100 次 系数 95% 区间 ==')
for t, f_ in enumerate(FEATS):
    v = sorted(x[t] for x in bw)
    print(f'  {f_:6s} [{v[2]:+.3f}, {v[97]:+.3f}]')
v = sorted(x[0]+x[1] for x in bw)
print(f'  极速+高度 [{v[2]:+.3f}, {v[97]:+.3f}]')

# ---- 全库预测, 低估/高估 ----
op = [r for r in usable if r['status'] == 'Operating']
for r in op:
    r['pred'] = sum(w[t]*x_of(r)[t] for t in range(len(FEATS)))
labeled_ids = set(id(e['row']) for e in steel + wood)

print('\n== 被低估候选: 预测分最高但不在 GTA 2018 任一 top50 ==')
unl = sorted([r for r in op if id(r) not in labeled_ids], key=lambda r: -r['pred'])
for r in unl[:15]:
    print(f"  {r['pred']:+.2f}  {r['name'][:30]:30s} {r['park'][:24]:24s} {r['country'][:14]:14s} {r['year']} "
          f"h={r['h']:.0f} s={r['s']:.0f} len={r['len']:.0f} inv={r['inv_i']:.0f} dur={r['dur_i']:.0f}{'*' if not r['dur'] else ''}")

print('\n== 被高估候选: 榜单名次远好于预测名次 (在榜车内比较) ==')
for lst, lbl in ((steel, '钢'), (wood, '木')):
    scored = sorted(lst, key=lambda e: -e['row']['pred'])
    predrank = {id(e['row']): i+1 for i, e in enumerate(scored)}
    resid = sorted(lst, key=lambda e: (e['rank'] - predrank[id(e['row'])]))
    for e in resid[:5]:
        r = e['row']
        print(f"  [{lbl}] GTA#{e['rank']:2d} vs 预测#{predrank[id(r)]:2d}  {r['name'][:28]:28s} "
              f"h={r['h']:.0f} s={r['s']:.0f} dur={r['dur_i']:.0f}{'*' if not r['dur'] else ''}")
