"""Cluster scene embeddings and report what the clusters actually contain. The question this answers: does a thematically coherent corpus contain RECURRING situations, or 800 singletons? A pack needs ~15 repeatable situations per stage; a cluster that draws from many different stories is a recurring situation, one that draws from a single story is just a scene. """ import json, math, pathlib, random, sys K = int(sys.argv[1]) if len(sys.argv) > 1 else 40 SEED = 11 chunks = json.loads(pathlib.Path('chunks.json').read_text(encoding='utf-8')) vecs = json.loads(pathlib.Path('embeddings.json').read_text(encoding='utf-8')) n = min(len(chunks), len(vecs)) chunks, vecs = chunks[:n], vecs[:n] print(f'clustering {n} chunks into k={K}\n') def norm(v): m = math.sqrt(sum(x * x for x in v)) or 1.0 return [x / m for x in v] vecs = [norm(v) for v in vecs] dim = len(vecs[0]) def dot(a, b): return sum(x * y for x, y in zip(a, b)) # k-means++ init on cosine distance (vectors are unit length, so dot == cos) random.seed(SEED) cent = [vecs[random.randrange(n)]] while len(cent) < K: d2 = [min((1 - dot(v, c)) for c in cent) ** 2 for v in vecs] tot = sum(d2) or 1.0 r = random.random() * tot acc = 0.0 for i, x in enumerate(d2): acc += x if acc >= r: cent.append(vecs[i]); break else: cent.append(vecs[random.randrange(n)]) assign = [0] * n for it in range(30): moved = 0 for i, v in enumerate(vecs): best, bs = 0, -2.0 for k, c in enumerate(cent): s = dot(v, c) if s > bs: bs, best = s, k if assign[i] != best: assign[i] = best; moved += 1 for k in range(K): mem = [vecs[i] for i in range(n) if assign[i] == k] if not mem: continue cent[k] = norm([sum(m[d] for m in mem) / len(mem) for d in range(dim)]) if moved == 0: break groups = {} for i, k in enumerate(assign): groups.setdefault(k, []).append(i) rows = [] for k, idx in groups.items(): stories = {chunks[i]['story'] for i in idx} coh = sum(dot(vecs[i], cent[k]) for i in idx) / len(idx) rows.append({'k': k, 'size': len(idx), 'stories': len(stories), 'coh': coh, 'idx': idx}) rows.sort(key=lambda r: (-r['stories'], -r['size'])) print(f"{'size':>5} {'stories':>8} {'coh':>6} sample titles") print('-' * 92) for r in rows: titles = [] for i in r['idx']: t = chunks[i]['title'] if t not in titles: titles.append(t) print(f"{r['size']:5} {r['stories']:8} {r['coh']:6.3f} {', '.join(titles[:4])[:66]}") multi = [r for r in rows if r['stories'] >= 4] print(f"\nclusters drawing on >=4 different stories: {len(multi)} of {K}") print(f"clusters that are essentially one story: {sum(1 for r in rows if r['stories'] <= 2)}") pathlib.Path('clusters.json').write_text(json.dumps( [{'k': r['k'], 'size': r['size'], 'stories': r['stories'], 'coh': r['coh'], 'idx': r['idx']} for r in rows]))