Download src/peppa/data.py from ChatterjeeLab/PepPA: direct link, hf CLI and curl.
- Browser
- Download file 1.91 kB
-
https://huggingface.co/ChatterjeeLab/PepPA/resolve/main/src/peppa/data.py
- Command line
-
hf download hf://ChatterjeeLab/PepPA/src/peppa/data.py
-
curl -L -o data.py https://huggingface.co/ChatterjeeLab/PepPA/resolve/main/src/peppa/data.py
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))] | |