# %% [code]
#!/usr/bin/env python3
"""
An inference example for CUDA-ported version of US-align (Zhang et al., 2022 Nature Methods).
Designed for binary stdin batch mode.
"""

import numpy as np, struct, subprocess, time, os, re, tempfile, zipfile
from collections import defaultdict

# ============================================================
# 1. COMPILE
# ============================================================
# Find dataset (Kaggle path varies)
for candidate in [
    '/kaggle/input/gpu-usalign-cuda',
    '/kaggle/input/datasets/gapchenko/gpu-usalign-cuda',
]:
    if os.path.exists(candidate):
        DATASET = candidate; break
else:
    raise RuntimeError('Dataset not found')

BUILD_DIR = '/tmp/gpu_usalign'
BINARY = os.path.join(BUILD_DIR, 'usalign_gpu')

print('=' * 60)
print('GPU USAlign Benchmark')
print('=' * 60)
print(f'Dataset: {DATASET}')
print(f'Contents: {os.listdir(DATASET)}')

os.makedirs(BUILD_DIR, exist_ok=True)

# Extract src.zip if needed, or copy src/ directly
src_zip = os.path.join(DATASET, 'src.zip')
src_dir = os.path.join(DATASET, 'src')
if os.path.exists(src_zip):
    print(f'Extracting {src_zip}...')
    with zipfile.ZipFile(src_zip, 'r') as z:
        z.extractall(BUILD_DIR)
    # Files might be in BUILD_DIR/src/ or BUILD_DIR/ depending on zip structure
    if os.path.exists(os.path.join(BUILD_DIR, 'src')):
        os.system(f'mv {BUILD_DIR}/src/* {BUILD_DIR}/')
elif os.path.isdir(src_dir):
    os.system(f'cp {src_dir}/* {BUILD_DIR}/')
else:
    # Try copying everything that looks like source
    for f in os.listdir(DATASET):
        if f.endswith(('.cu', '.cuh', '.py')) or f == 'Makefile':
            os.system(f'cp {DATASET}/{f} {BUILD_DIR}/')

print(f'Build dir: {os.listdir(BUILD_DIR)[:10]}...')

# Detect GPU arch
try:
    cap = subprocess.check_output(
        'nvidia-smi --query-gpu=compute_cap --format=csv,noheader',
        shell=True, text=True).strip().split('\n')[0]
    arch = 'sm_' + cap.replace('.', '')
except Exception:
    arch = 'sm_75'

print(f'\nCompiling for {arch}...')
t0 = time.time()
ret = subprocess.run(
    f'cd {BUILD_DIR} && /usr/local/cuda/bin/nvcc -O3 -std=c++17 '
    f'-Xcompiler -O3,-ffast-math -arch={arch} '
    f'-o usalign_gpu usalign_gpu.cu 2>&1',
    shell=True, capture_output=True, text=True)
print(ret.stdout[-500:] if ret.stdout else '')
if ret.stderr:
    # Show only errors, not warnings
    for line in ret.stderr.split('\n'):
        if 'error' in line.lower():
            print(line)
print(f'Compiled in {time.time()-t0:.1f}s')

if not os.path.exists(BINARY):
    print(f'FULL STDERR:\n{ret.stderr}')
    raise RuntimeError('Compilation failed!')

# CPU USalign — from dataset or system
CPU_USALIGN = None
for p in [os.path.join(BUILD_DIR, 'USalign'),
          '/usr/bin/USalign', '/opt/conda/bin/USalign', '/usr/local/bin/USalign']:
    if os.path.exists(p):
        os.chmod(p, 0o755)
        CPU_USALIGN = p; break
if CPU_USALIGN is None:
    os.system('pip install -q usalign 2>/dev/null')
    CPU_USALIGN = 'USalign'
print(f'CPU USalign: {CPU_USALIGN}')

# ============================================================
# 2. LOAD BENCHMARK PAIRS
# ============================================================
print(f'\n{"="*60}')
print('Loading benchmark pairs (200 real RNA from training data)')
print(f'{"="*60}')

meta = np.load(os.path.join(DATASET, 'benchmark_pairs.npz'), allow_pickle=True)
cdata = np.load(os.path.join(DATASET, 'benchmark_pairs_coords.npz'))

tids1, tids2 = meta['tids1'], meta['tids2']
lens1, lens2 = meta['lens1'], meta['lens2']
seqs1, seqs2 = meta['seqs1'], meta['seqs2']
n_pairs = int(meta['n_pairs'])
off1, off2 = cdata['offsets1'], cdata['offsets2']
c1_list = [cdata['coords1'][off1[i]:off1[i+1]] for i in range(n_pairs)]
c2_list = [cdata['coords2'][off2[i]:off2[i+1]] for i in range(n_pairs)]

print(f'  {n_pairs} pairs')
for lo, hi in [(20,50),(50,100),(100,300),(300,800),(800,3000)]:
    cnt = sum(lo <= max(l1,l2) < hi for l1,l2 in zip(lens1, lens2))
    if cnt: print(f'    [{lo:>4}-{hi:<4}): {cnt} pairs')

# ============================================================
# 3. GPU BENCHMARK (batched by size bucket)
# ============================================================
print(f'\n{"="*60}')
print('GPU USAlign — batch mode (binary stdin, grouped by length)')
print(f'{"="*60}')

buckets = defaultdict(list)
for i in range(n_pairs):
    ml = max(len(c1_list[i]), len(c2_list[i]))
    for lo, hi in [(20,50),(50,100),(100,300),(300,800),(800,3000)]:
        if lo <= ml < hi: buckets[(lo,hi)].append(i); break

gpu_results = [None] * n_pairs
gpu_bucket_times = {}
gpu_total = 0

for (lo, hi), indices in sorted(buckets.items()):
    n = len(indices)
    buf = struct.pack('<II', n, 1)
    for i in indices:
        buf += struct.pack('<II', len(c1_list[i]), len(c2_list[i]))
        buf += c1_list[i].astype(np.float64).tobytes()
        buf += c2_list[i].astype(np.float64).tobytes()

    t0 = time.time()
    proc = subprocess.run([BINARY, '-stdin', '-mol', '1'],
                          input=buf, capture_output=True, timeout=600)
    elapsed = time.time() - t0
    gpu_total += elapsed
    gpu_bucket_times[(lo,hi)] = (elapsed, n)

    out = proc.stdout; off = 0
    for k, idx in enumerate(indices):
        tm1, tm2, rmsd = struct.unpack_from('<ddd', out, off); off += 24
        n8 = struct.unpack_from('<I', out, off)[0]; off += 4
        gpu_results[idx] = {'TM1': tm1, 'TM2': tm2, 'rmsd': rmsd, 'n_ali8': n8}

    print(f'  [{lo:>4}-{hi:<4}): {n:>3} pairs in {elapsed:.2f}s ({elapsed/n*1000:.1f} ms/pair)')

print(f'  TOTAL: {n_pairs} pairs in {gpu_total:.2f}s ({gpu_total/n_pairs*1000:.1f} ms/pair)')

# ============================================================
# 4. CPU BENCHMARK
# ============================================================
print(f'\n{"="*60}')
print('CPU USAlign — sequential (PDB files)')
print(f'{"="*60}')

tmpdir = tempfile.mkdtemp()
cpu_results = {}
cpu_total = 0
MAX_CPU = 240

def write_pdb(coords, seq, path):
    rn = {'A': '  A', 'U': '  U', 'G': '  G', 'C': '  C'}
    with open(path, 'w') as f:
        for i in range(len(coords)):
            r = rn.get(seq[i].upper(), '  N') if i < len(seq) else '  N'
            f.write(f'ATOM  {i+1:5d}  C1\' {r} A{i+1:4d}    '
                    f'{coords[i,0]:8.3f}{coords[i,1]:8.3f}{coords[i,2]:8.3f}'
                    f'  1.00  0.00           C\n')

for idx in range(n_pairs):
    if cpu_total > MAX_CPU:
        print(f'  Stopped at {idx} pairs (time limit {MAX_CPU}s)')
        break
    write_pdb(c1_list[idx], str(seqs1[idx]), f'{tmpdir}/a.pdb')
    write_pdb(c2_list[idx], str(seqs2[idx]), f'{tmpdir}/b.pdb')
    t0 = time.time()
    proc = subprocess.run([CPU_USALIGN, f'{tmpdir}/a.pdb', f'{tmpdir}/b.pdb', '-atom', " C1'"],
                          capture_output=True, text=True, timeout=120)
    ms = (time.time() - t0) * 1000
    cpu_total += ms / 1000
    m = re.search(r'TM-score=\s*([\d.]+)\s*\(normalized by length of Structure_1', proc.stdout)
    cpu_results[idx] = {'TM2': float(m.group(1)) if m else -1, 'ms': ms}
    if (idx+1) % 50 == 0:
        print(f'  {idx+1} pairs, {cpu_total:.1f}s')

n_cpu = len(cpu_results)
print(f'  {n_cpu} pairs in {cpu_total:.1f}s ({cpu_total/max(n_cpu,1)*1000:.0f} ms/pair)')

# ============================================================
# 5. RESULTS
# ============================================================
gpu_tms = np.array([gpu_results[i]['TM2'] for i in cpu_results])
cpu_tms = np.array([cpu_results[i]['TM2'] for i in cpu_results])
valid = cpu_tms >= 0
diffs = gpu_tms[valid] - cpu_tms[valid]
corr = np.corrcoef(gpu_tms[valid], cpu_tms[valid])[0,1]

print(f'\n{"="*60}')
print('ACCURACY')
print(f'{"="*60}')
print(f'  Pairs:        {valid.sum()}')
print(f'  Correlation:  {corr:.4f}')
print(f'  Mean |delta|: {np.mean(np.abs(diffs)):.5f}')
print(f'  Max  |delta|: {np.max(np.abs(diffs)):.5f}')
print(f'  Within 1%:    {np.mean(np.abs(diffs)<0.01)*100:.0f}%')
print(f'  Within 3%:    {np.mean(np.abs(diffs)<0.03)*100:.0f}%')

gpu_avg = gpu_total / n_pairs * 1000
cpu_avg = cpu_total / max(n_cpu, 1) * 1000

print(f'\n{"="*60}')
print('SPEED')
print(f'{"="*60}')
print(f'  GPU batch:    {gpu_avg:.1f} ms/pair')
print(f'  CPU seq:      {cpu_avg:.0f} ms/pair')
print(f'  Speedup:      {cpu_avg/max(gpu_avg,0.1):.1f}x')

print(f'\n  {"Bucket":<15} {"GPU ms":>8} {"CPU ms":>8} {"Speedup":>8}')
print(f'  {"-"*43}')
for lo, hi in [(20,50),(50,100),(100,300),(300,800),(800,3000)]:
    cm = [cpu_results[i]['ms'] for i in cpu_results if lo <= max(lens1[i],lens2[i]) < hi]
    if not cm: continue
    ac = np.mean(cm)
    if (lo,hi) in gpu_bucket_times:
        gt, gn = gpu_bucket_times[(lo,hi)]
        ag = gt/gn*1000
    else: ag = gpu_avg
    print(f'  [{lo:>4}-{hi:<4})    {ag:>7.1f}  {ac:>7.0f}  {ac/max(ag,0.1):>7.1f}x')

print(f'\n{"="*60}')
print('DONE')
print(f'{"="*60}')
print(f'  Accuracy: corr={corr:.3f}, mean |delta|={np.mean(np.abs(diffs)):.4f}')
print(f'  Speed: {cpu_avg/max(gpu_avg,0.1):.0f}x faster (batch mode)')
print(f'  Source: {DATASET}/src/')
