CIS6270 / lecture_7 /discrete_learning.py
pranamanam's picture
Add Lecture 7 OT and Schrödinger bridge implementations and equation readings
d8f9639 verified
Raw History Blame Contribute Delete
3.83 kB
"""Learn the discrete Markov projections used inside DDSBM and CSBM.
Run: python discrete_learning.py --output results
The outer exact IMF loop is in ot_sbm_examples.py; this file tests learned
rate/transition fitting at its known small bridge solution.
"""
import argparse,json,math
from pathlib import Path
import numpy as np
import torch
from torch import nn
from torch.nn import functional as F
from ot_sbm_examples import ctmc_bridge_example,finite_bridge_example
from learned_bridges import save_artifact
torch.set_num_threads(1)
def kernel(dt):
z=torch.exp(-3*dt)
return torch.stack([1/3+2*z/3,2*(1-z)/3,(1-z)/3,2/3+z/3],-1).reshape(-1,2,2)
def rate_model():
return nn.Sequential(nn.Linear(3,32),nn.SiLU(),nn.Linear(32,32),nn.SiLU(),nn.Linear(32,1),nn.Softplus())
def fit_ddsbm(seed=13,updates=1400,batch=512,checkpoint=None):
torch.manual_seed(seed);exact=ctmc_bridge_example();pi=torch.tensor(exact['coupling'].reshape(-1),dtype=torch.float32)
K=kernel(torch.tensor([1.]))[0];Goff=torch.tensor([2.,1.]);model=rate_model();opt=torch.optim.Adam(model.parameters(),lr=.002)
for _ in range(updates):
pair=torch.multinomial(pi,batch,replacement=True);a=pair//2;z=pair%2;t=.02+.96*torch.rand(batch)
left=kernel(t);right=kernel(1-t);idx=torch.arange(batch)
# P(X_t=x | X_0=a, X_1=z) under the reference bridge.
bridge=left[idx,a,:]*right[idx,:,z]/K[a,z,None]
x=torch.multinomial(bridge,1).squeeze(1);y=1-x
target=Goff[x]*right[idx,y,z]/right[idx,x,z]
q=model(torch.cat([t[:,None],F.one_hot(x,2).float()],1)).squeeze(1)
loss=(q-target*q.clamp_min(1e-8).log()).mean()
opt.zero_grad();loss.backward();opt.step()
# Compare predicted rates with the independently computed Doob transform.
truth=[];pred=[]
with torch.no_grad():
for row in exact['values']:
t=float(row['t'])
for x in [0,1]:
pred.append(float(model(torch.tensor([[t,float(x==0),float(x==1)]]))))
truth.append(float(row['rates'][x,1-x]))
rel=np.mean(((np.array(pred)-truth)/truth)**2)
save_artifact(checkpoint,method='ddsbm',model=model.state_dict(),steps=100)
return dict(predicted_off_diagonal_rates=pred,exact_off_diagonal_rates=truth,relative_mse=float(rel),
approximation='Neural Markov rate projection for a known finite-state SB coupling; outer IMF is separately enumerated exactly.')
def fit_csbm(seed=14,updates=600,checkpoint=None):
torch.manual_seed(seed);ex=finite_bridge_example(steps=2);paths=torch.tensor(ex['paths']);mass=torch.tensor(ex['bridge'],dtype=torch.float32)
logits=nn.Parameter(torch.zeros(2,2,2));opt=torch.optim.Adam([logits],lr=.04)
for _ in range(updates):
loss=torch.tensor(0.)
for k in range(2):
logq=logits[k].log_softmax(-1)
loss-=torch.sum(mass*logq[paths[:,k],paths[:,k+1]])
opt.zero_grad();loss.backward();opt.step()
q=logits.softmax(-1).detach().numpy();truth=np.array(ex['Q']);p=np.array([.5,.5])
for transition in q:p=p@transition
save_artifact(checkpoint,method='csbm',transitions=torch.tensor(q),steps=2)
return dict(transitions=q.tolist(),maximum_transition_error=float(abs(q-truth).max()),terminal=p.tolist(),
approximation='Exact path-weighted categorical cross-entropy with a fully expressive transition table; no sampling error in this fit.')
def main():
parser=argparse.ArgumentParser();parser.add_argument('--output',default='results');args=parser.parse_args();out=Path(args.output);out.mkdir(parents=True,exist_ok=True)
result={'ddsbm':fit_ddsbm(),'csbm':fit_csbm()};(out/'discrete_learned_results.json').write_text(json.dumps(result,indent=2));print(json.dumps(result,indent=2))
if __name__=='__main__':main()