PepPA / src /peppa /compiled.py
pranamanam's picture
Upload 97 files
98bde72 verified
Raw
History Blame Contribute Delete
10.8 kB
"""Deterministic reference kernel for frozen computational plans.
Scientific workers, semantic evidence adjudication and chemical standardization
are external contracts. No experimental result action exists in this kernel.
"""
from __future__ import annotations
from copy import deepcopy
from dataclasses import dataclass
from typing import Callable
import math
import numpy as np
from scipy.optimize import minimize
from .schema import canonical, digest
@dataclass(frozen=True)
class Node:
id: str
tool: str
revision: str
inputs: tuple[str, ...]
outputs: tuple[str, ...]
cost: int = 1
attempts: int = 1
@dataclass(frozen=True)
class Plan:
nodes: tuple[Node, ...]
required: tuple[str, ...]
budget: int
seed: int = 2027
def validate_plan(plan: Plan, initial: dict, registry: dict) -> tuple[Node, ...]:
"""Validate unique producers, complete dependencies, DAG, revisions and cap."""
if plan.budget < 0 or not plan.required or len(set(plan.required)) != len(plan.required):
raise ValueError('invalid budget or required artifacts')
ids, produced = set(), set(initial)
for n in plan.nodes:
if not n.id or n.id in ids or n.tool not in registry:
raise ValueError('duplicate node or unknown tool')
if n.revision != registry[n.tool]['revision']:
raise ValueError('revision mismatch')
if n.cost < 1 or not 1 <= n.attempts <= 3 or not n.outputs:
raise ValueError('invalid resource contract')
if len(set(n.outputs)) != len(n.outputs) or produced.intersection(n.outputs):
raise ValueError('artifact overwrite')
produced.update(n.outputs); ids.add(n.id)
if not set(plan.required) <= produced:
raise ValueError('uncovered release obligation')
if sum(n.cost*n.attempts for n in plan.nodes) > plan.budget:
raise ValueError('worst-case budget exceeded')
ordered, pending, available = [], list(plan.nodes), set(initial)
while pending:
ready = sorted((n for n in pending if set(n.inputs) <= available), key=lambda n:n.id)
if not ready:
raise ValueError('cycle or missing input')
n = ready[0]; ordered.append(n); available.update(n.outputs); pending.remove(n)
return tuple(ordered)
def execute(plan: Plan, initial: dict, registry: dict, cache: dict | None = None) -> dict:
"""Atomic artifact commits; deterministic schedule and hash-linked records.
Registry values: revision, run(payload, seed), validate(output). Validation
must reject wrong units, chemistry, identities and out-of-domain use.
Operational timestamps belong in a separate log, outside semantic hashes.
"""
order = validate_plan(plan, initial, registry)
artifacts = deepcopy(initial); events=[]; spent=0
cache = {} if cache is None else cache
def append(payload):
event={'index':len(events),'previous':events[-1]['hash'] if events else '0'*64, **payload}
event['hash']=digest(event); events.append(event)
plan_data={'nodes':[vars(n) for n in order], 'required':plan.required,'budget':plan.budget,'seed':plan.seed}
append({'kind':'initialize','plan_hash':digest(plan_data),'initial':deepcopy(initial)})
for n in order:
if not set(n.inputs) <= set(artifacts):
append({'kind':'blocked','node':n.id,'reason':'DEPENDENCY_FAILED'});continue
payload={k:deepcopy(artifacts[k]) for k in n.inputs}
request={'node':vars(n),'inputs':payload,'seed':plan.seed}
key=digest(request)
for attempt in range(n.attempts):
spent+=n.cost
try:
if key in cache:
saved=cache[key]
if saved['hash'] != digest(saved['value']):raise ValueError('CACHE_CORRUPT')
result=deepcopy(saved['value'])
else:
result=registry[n.tool]['run'](deepcopy(payload),plan.seed)
if not isinstance(result,dict) or set(result)!=set(n.outputs):
raise ValueError('OUTPUT_CONTRACT')
canonical(result) # rejects NaN/infinity and unserializable values
registry[n.tool]['validate'](result)
cache[key]={'value':deepcopy(result),'hash':digest(result)}
artifacts.update(deepcopy(result))
append({'kind':'commit','node':n.id,'request':key,'attempt':attempt+1,'cost':n.cost,'result':deepcopy(result)})
break
except Exception as e:
append({'kind':'failure','node':n.id,'request':key,'attempt':attempt+1,'cost':n.cost,'error':type(e).__name__+': '+str(e)})
missing=sorted(set(plan.required)-set(artifacts))
append({'kind':'close','spent':spent,'missing':missing,'status':'complete' if not missing else 'incomplete'})
return {'artifacts':artifacts,'events':events,'complete':not missing,'spent':spent}
def replay(events: list[dict]) -> dict:
artifacts={};prev='0'*64
if not events or events[0].get('kind')!='initialize' or events[-1].get('kind')!='close':
raise ValueError('incomplete trace')
for i,item in enumerate(events):
e=deepcopy(item); h=e.pop('hash')
if e['index']!=i or e['previous']!=prev or digest(e)!=h:raise ValueError('trace corruption')
prev=h
if e['kind']=='initialize':artifacts=deepcopy(e['initial'])
if e['kind']=='commit':
if set(artifacts).intersection(e['result']):raise ValueError('artifact overwrite')
artifacts.update(deepcopy(e['result']))
return artifacts
def collapse_lineage(values, lineages):
"""Repeated use of an identical model output has exactly one ensemble vote."""
a=np.asarray(values,dtype=float)
if a.ndim != 2 or a.shape[1]!=len(lineages) or not np.isfinite(a).all():
raise ValueError('invalid predictions')
groups={}
for j,k in enumerate(lineages):groups.setdefault(k,[]).append(j)
cols=[]
for k in sorted(groups):
js=groups[k]
if not all(np.array_equal(a[:,js[0]],a[:,j]) for j in js):
raise ValueError('inconsistent duplicate lineage')
cols.append(a[:,js[0]])
return np.column_stack(cols), sorted(groups)
def fit_ensemble(predictions, y, lineages, ridge=0.01):
"""Nonnegative simplex stacking on a fit partition, before calibration."""
a,keys=collapse_lineage(predictions,lineages); y=np.asarray(y,float)
if len(a)==0 or y.shape!=(len(a),) or not np.isfinite(y).all() or ridge<=0:
raise ValueError('invalid fit partition')
m=a.shape[1];w0=np.full(m,1/m)
def loss(w):return np.mean((a@w-y)**2)+ridge*np.sum((w-w0)**2)
r=minimize(loss,w0,method='SLSQP',bounds=[(0,1)]*m,constraints={'type':'eq','fun':lambda w:w.sum()-1},options={'ftol':1e-12,'maxiter':1000})
if not r.success:raise ValueError(r.message)
w=np.maximum(r.x,0);w/=w.sum()
return {'lineages':keys,'weights':w.tolist(),'ridge':ridge}
def aggregate(predictions,lineages,fit):
a,keys=collapse_lineage(predictions,lineages)
if keys!=fit['lineages']:raise ValueError('ensemble membership changed')
w=np.asarray(fit['weights']);mean=a@w
return mean,np.sqrt(np.sum(w*(a-mean[:,None])**2,axis=1))
def release(candidates: list[dict], obligations: list[dict], k=12) -> dict:
"""Complete computational records with typed prediction/proxy support.
Required per score: value, lower, upper, unit, support, lineage, domain_ok.
Proxy bounds are ensemble extrema; calibrated bounds are statistical.
Caller provides standardized chemistry identities after scientific preflight.
"""
from .metrics import diverse_select
if not obligations or k<0 or len({o['id'] for o in obligations})!=len(obligations):
raise ValueError('invalid obligations')
records=[]; seen=set()
for c in sorted(deepcopy(candidates),key=lambda x:x['id']):
if c['id'] in seen:raise ValueError('duplicate chemical identity')
seen.add(c['id']);errors=[];margins=[]
if not c.get('chemistry_verified'):errors.append('CHEMISTRY_UNVERIFIED')
for o in obligations:
s=c.get('scores',{}).get(o['id'])
if s is None:errors.append(o['id']+':MISSING');continue
try:
vals=[s[x] for x in ['value','lower','upper']]
if not all(isinstance(v,(int,float)) and math.isfinite(v) for v in vals):raise ValueError()
if not s['lower']<=s['value']<=s['upper']:raise ValueError()
if s['unit']!=o['unit'] or s['support']!=o['support'] or not s['lineage'] or not s['domain_ok']:raise ValueError()
if o['scale']<=0 or o['direction'] not in ['ge','le']:raise ValueError()
v=s['lower'] if o['direction']=='ge' else s['upper']
margin=(v-o['threshold'])/o['scale']*(1 if o['direction']=='ge' else -1)
margins.append(margin)
if margin<0:errors.append(o['id']+':THRESHOLD')
except (KeyError,TypeError,ValueError):errors.append(o['id']+':INVALID')
c.update(errors=errors,eligible=not errors,worst_margin=min(margins) if margins else None)
records.append(c)
scores={c['id']:c['worst_margin'] for c in records if c['eligible']}
sequences={c['id']:c.get('monomers',c['sequence']) for c in records}
chosen=diverse_select(sequences,scores,k)
return {'records':records,'selected':chosen,'unfilled':k-len(chosen),'obligations_hash':digest(obligations)}
def robust_reward(margins, weights):
"""Rows=endpoints, columns=required contexts; manuscript Eq. robustgen."""
m=np.asarray(margins,float);w=np.asarray(weights,float)
if m.ndim!=2 or not m.shape[1] or w.shape!=(m.shape[0],) or not np.isfinite(m).all() or not np.isfinite(w).all() or (w<0).any() or not np.isclose(w.sum(),1):
raise ValueError('invalid margins or weights')
u=m.min(axis=1)
return float(-np.max(w*np.maximum(0,-u))+.05*np.sum(w*np.tanh(u)))
def redesign_schedule(failures: dict[str,int], slots: int) -> list[str]:
"""Smoothed failure-proportional quotas, largest remainder, stable scheduling."""
if slots<0 or not failures or any(not isinstance(v,int) or v<0 for v in failures.values()):
raise ValueError('invalid failure counts')
order=sorted(failures,key=lambda k:(-failures[k],k))
denominator=sum(failures.values())+len(failures)
exact={k:slots*(failures[k]+1)/denominator for k in order}
quota={k:math.floor(v) for k,v in exact.items()}
residual=slots-sum(quota.values())
for k in sorted(order,key=lambda k:(-(exact[k]-quota[k]),k))[:residual]:quota[k]+=1
schedule=[]
while len(schedule)<slots:
for k in order:
if quota[k]>0:schedule.append(k);quota[k]-=1
return schedule