PepPA / src /peppa /data.py
pranamanam's picture
Upload 97 files
98bde72 verified
Raw
History Blame Contribute Delete
1.91 kB
"""Label-preserving dataset normalization and connected-component splitting."""
from collections import defaultdict
import hashlib,json
LABELS={'positive':1,'negative':0,'unknown':None,'pos':1,'neg':0,'1':1,'0':0}
def normalize_snooppi(row):
raw=str(row['SNOOPPI_final_label']).lower().strip()
if raw not in LABELS:raise ValueError('unrecognized SNOOPPI label: '+raw)
return {'label':LABELS[raw],'label_source':raw,'partner_a':row['partner_A_sequence'],
'partner_b':row['partner_B_sequence'],'source':row.get('pubmed_ids',''),
'publication_years':row.get('publication_years',''),'assays':row.get('assays','')}
def connected_splits(rows,seed=2027):
"""Union target clusters, peptide clusters, and publication IDs across rows.
Cluster IDs are computed before this function (e.g. MMseqs2). All rows
sharing any entity or primary publication receive the same split.
"""
parent=list(range(len(rows)))
def root(i):
while parent[i]!=i:parent[i]=parent[parent[i]];i=parent[i]
return i
def union(i,j):
a,b=root(i),root(j)
if a!=b:parent[max(a,b)]=min(a,b)
seen={}
for i,r in enumerate(rows):
groups=['target:'+v for v in r['target_clusters']]+['peptide:'+v for v in r['peptide_clusters']]+['source:'+v for v in r['source_ids']]
if not groups:raise ValueError('missing split source records')
for g in groups:
if g in seen:union(i,seen[g])
else:seen[g]=i
members=defaultdict(list)
for i,r in enumerate(rows):members[root(i)].append(str(r['id']))
labels={}
for k,ids in members.items():
h=int(hashlib.sha256((str(seed)+'|'+ '|'.join(sorted(ids))).encode()).hexdigest()[:12],16)/16**12
labels[k]='train' if h<.7 else 'ensemble_fit' if h<.8 else 'calibration' if h<.9 else 'test'
return [labels[root(i)] for i in range(len(rows))]