CIS6270 / lecture_6 /run.py
pranamanam's picture
Add Lecture 6 flow maps with tested training, sampling, and course navigation
da3babb verified
Raw History Blame Contribute Delete
14.1 kB
#!/usr/bin/env python3
"""Train, save, reload, and sample the Lecture 6 flow-map methods."""
import argparse
import json
import platform
import time
from pathlib import Path
import torch
from torch.nn import functional as F
from common import (MapNet, SequenceNet, interpolate, load_text, mixture_data,
mixture_metrics, seed_all, text_metrics, write_json)
from continuous import (Autoencoder, embed_surface, integrate_velocity, sample_continuous,
train_continuous, train_latent)
from categorical import sample_categorical, train_categorical
from posterior import (guided_samples, posterior_diagnostics, posterior_samples,
posterior_value, reward, train_posterior, train_reward_drift,
weighted_diamond_samples)
from expanding import ExpandingNet, sample_expanding, train_expanding
from stochastic import StrongMap, evaluate_stochastic, train_stochastic
ROOT = Path(__file__).resolve().parent
CONTINUOUS = ['flow-matching','fmm-lagrangian','fmm-eulerian','self-distill',
'consistency','shortcut','meanflow','latent']
CATEGORICAL = ['fmlm','categorical','discrete-lsd','discrete-esd']
METHODS = CONTINUOUS+CATEGORICAL+['diamond','meta','expanding','ssfm']
def parser():
p=argparse.ArgumentParser(description=__doc__)
p.add_argument('--method',choices=METHODS)
p.add_argument('--mode',choices=['train-sample','train','sample'],default='train-sample')
p.add_argument('--train-steps',type=int,default=1000)
p.add_argument('--teacher-steps',type=int,default=1000)
p.add_argument('--finetune-steps',type=int,default=0,help='Optional Meta reward-drift fine-tuning')
p.add_argument('--sample-steps',type=int,default=8)
p.add_argument('--posterior-steps',type=int,default=4)
p.add_argument('--particles',type=int,default=32)
p.add_argument('--reward-strength',type=float,default=1.)
p.add_argument('--ssfm-target',choices=['official-code','paper'],default='official-code')
p.add_argument('--batch-size',type=int,default=64)
p.add_argument('--samples',type=int,default=128)
p.add_argument('--width',type=int,default=64)
p.add_argument('--lr',type=float,default=1e-3)
p.add_argument('--seed',type=int,default=6270)
p.add_argument('--threads',type=int,default=1)
p.add_argument('--device',choices=['cpu','cuda'],default='cpu')
p.add_argument('--data',help='Whitespace-separated text, one sequence per line; or continuous CSV with two numeric columns')
p.add_argument('--out',help='Run directory; defaults to lecture_6/outputs/METHOD')
return p
def restore_model(checkpoint,device):
method=checkpoint['method'];width=checkpoint['width']
if method in CATEGORICAL:
model=SequenceNet(checkpoint['length'],len(checkpoint['vocab']),width)
elif method=='expanding':
model=ExpandingNet(checkpoint['length'],len(checkpoint['vocab']),width)
elif method=='ssfm':
model=StrongMap(width,checkpoint['sigma'])
else:
model=MapNet(checkpoint['dim'],width,3 if method in ['diamond','meta'] else 0)
model.load_state_dict(checkpoint['model'])
model=model.to(device).eval()
return model
def generate(model,state,args):
"""All metrics use fresh, seeded samples; no training examples stand in for output."""
seed_all(args.seed+100,args.threads)
method=state['method'];device=args.device
report={'method':method,'sample_steps':args.sample_steps,'samples':args.samples,
'sampling_seed':args.seed+100}
if method in CATEGORICAL:
ids=sample_categorical(model,args.samples,args.sample_steps,device).cpu()
texts,metrics=text_metrics(ids,state['vocab'],reference=set(state['reference_text']))
report.update(metrics)
return texts,report
if method=='expanding':
ids,lengths,extra=sample_expanding(model,args.samples,args.sample_steps,device)
texts,metrics=text_metrics(ids.cpu(),state['vocab'],lengths.cpu(),set(state['reference_text']))
report.update(metrics);report.update(extra)
return texts,report
if method=='ssfm':
x,extra=evaluate_stochastic(model,args.samples,args.sample_steps,device)
report.update(extra)
elif method in ['diamond','meta']:
with torch.no_grad():
outer=torch.randn(args.samples,2,device=device)
t=torch.zeros(args.samples,1,device=device)
x=posterior_samples(model,outer,t,1,args.posterior_steps)[:,0]
guided=guided_samples(model,args.samples,args.sample_steps,args.particles,
args.reward_strength,args.posterior_steps)
report.update(posterior_diagnostics(model,args.posterior_steps))
report['unconditional_mean_reward']=float(reward(x,args.reward_strength).mean())
report['guided_mean_reward']=float(reward(guided,args.reward_strength).mean())
report['guided_distribution']=mixture_metrics(guided)
report['posterior_steps']=args.posterior_steps
report['particles']=args.particles
observation=torch.tensor([[.5,-.3]],device=device)
_,estimate,ess=weighted_diamond_samples(observation,.5,4096)
report['importance_posterior_mean']=estimate.tolist()
report['importance_effective_sample_size']=ess.tolist()
# Verify differentiability through the learned posterior's context.
observation=observation.detach().requires_grad_(True)
value,_,_=posterior_value(model,observation,.5,args.particles,args.reward_strength,args.posterior_steps)
gradient=torch.autograd.grad(value.sum(),observation)[0]
report['posterior_value_context_gradient']=gradient.tolist()
report['finite_value_gradient']=bool(torch.isfinite(gradient).all())
report['guided_samples']=guided.detach().cpu().tolist()
if 'reward_drift' in state:
drift=MapNet(2,state['width']).to(device)
drift.load_state_dict(state['reward_drift'])
# A complete sampler for the fitted drift; boundary-time queries
# extrapolate beyond the [.05,.85] fine-tuning interval.
aligned=integrate_velocity(lambda z,a:drift(z,a,a),outer,0.,1.,args.sample_steps)
report['finetuned_mean_reward']=float(reward(aligned,args.reward_strength).mean())
report['finetuned_finite_samples']=bool(torch.isfinite(aligned).all())
report['finetuned_samples']=aligned.detach().cpu().tolist()
report['finetuned_sampling_note']='Full-interval Heun integration; boundary times extrapolate beyond fine-tuning support.'
else:
noise=torch.randn(args.samples,state['dim'],device=device)
x=sample_continuous(model,method,noise,args.sample_steps)
if method=='latent':
ae=Autoencoder(state['width']).to(device)
ae.load_state_dict(state['autoencoder'])
with torch.no_grad():
x=ae.decoder(x*state['latent_std'].to(device)+state['latent_mean'].to(device))
heldout=state['heldout_data'].to(device)
reconstruction=ae.decoder(ae.encoder(embed_surface(heldout)))
report['heldout_reconstruction_mse']=float((reconstruction-embed_surface(heldout)).square().mean())
report['surface_residual_rmse']=float((x[:,2:]-.3*(x[:,:1].square()-x[:,1:2].square())).square().mean().sqrt())
if method not in ['consistency','meanflow','flow-matching']:
from common import finite_map
with torch.no_grad():
direct=finite_map(model,noise,0.,1.)
split=finite_map(model,finite_map(model,noise,0.,.5),.5,1.)
report['composition_rmse']=float((direct-split).square().mean().sqrt())
report['finite_samples']=bool(torch.isfinite(x).all())
if not report['finite_samples']:raise FloatingPointError('Sampling produced nonfinite coordinates.')
if x.shape[1]>=2 and state['data_kind']=='four-gaussian-mixture':
report.update(mixture_metrics(x[:,:2]))
rows=[' '.join(f'{value:.7f}' for value in row) for row in x.detach().cpu().tolist()]
return rows,report
def main(argv=None):
p=parser();args=p.parse_args(argv)
if min(args.train_steps,args.teacher_steps,args.sample_steps,args.posterior_steps,
args.particles,args.batch_size,args.samples,args.width,args.threads)<1 or args.finetune_steps<0:
p.error('Step counts, sizes, and width must be positive; finetune steps must be nonnegative.')
if args.lr<=0:p.error('Learning rate must be positive.')
if args.device=='cuda' and not torch.cuda.is_available():p.error('CUDA is unavailable; select cpu.')
if args.mode=='sample' and args.out is None:p.error('Sample mode requires --out for its checkpoint.')
if args.mode!='sample' and args.method is None:args.method='flow-matching'
out=Path(args.out) if args.out else ROOT/'outputs'/args.method
out.mkdir(parents=True,exist_ok=True)
seed_all(args.seed,args.threads)
start=time.perf_counter()
if args.mode=='sample':
state=torch.load(out/'checkpoint.pt',map_location=args.device,weights_only=True)
if args.method is not None and args.method!=state['method']:
p.error('Requested method differs from the saved checkpoint.')
args.method=state['method']
if args.method=='ssfm' and args.sample_steps & (args.sample_steps-1):
p.error('SSFM sample steps must be a power of two.')
model=restore_model(state,args.device)
rows,report=generate(model,state,args)
(out/'resampled.txt').write_text('\n'.join(rows)+'\n')
write_json(out/'sample_report.json',report)
print(json.dumps({'method':args.method,'mode':'sample','out':str(out),'seconds':time.perf_counter()-start}))
return report
if args.method=='ssfm' and args.sample_steps & (args.sample_steps-1):
p.error('SSFM sample steps must be a power of two.')
if args.finetune_steps and args.method!='meta':p.error('--finetune-steps applies to meta.')
if args.data and args.method in ['diamond','meta','ssfm']:
p.error('Analytic posterior/OU examples use their specified reference distributions.')
stage_logs=[]
if args.method in CATEGORICAL+['expanding']:
default='variable_text.txt' if args.method=='expanding' else 'phrases.txt'
data_path=Path(args.data) if args.data else ROOT/'data'/default
ids,lengths,vocab=load_text(data_path,args.method=='expanding')
perm=torch.randperm(len(ids));ids,lengths=ids[perm],lengths[perm]
split=max(1,int(.8*len(ids)))
training=ids[:split].to(args.device)
if args.method=='expanding':
model,state,logs=train_expanding(training,lengths[:split].to(args.device),vocab,args)
else:
model,state,logs=train_categorical(args.method,training,vocab,args)
reference,_=text_metrics(ids[:split],vocab,lengths[:split])
state.update({'reference_text':reference,'data_kind':'synthetic-text' if args.data is None else 'custom-text'})
with torch.no_grad():
if args.method in CATEGORICAL:
heldout=ids[split:].to(args.device)
clean=F.one_hot(heldout,len(vocab)).float();t=torch.full((len(clean),1),.5,device=args.device)
x,_=interpolate(clean,t)
state['validation_ce']=float(F.cross_entropy(model(x,t,t).flatten(0,1),heldout.flatten()))
elif args.method=='ssfm':
model,state,logs=train_stochastic(args,args.device)
state['data_kind']='ornstein-uhlenbeck'
else:
if args.data:
import numpy as np
data=torch.tensor(np.loadtxt(args.data,delimiter=','),dtype=torch.float32)
if data.ndim!=2 or data.shape[1]!=2 or len(data)<8 or not torch.isfinite(data).all():
p.error('Continuous CSV needs at least eight finite rows and exactly two columns, without a header.')
data=data[torch.randperm(len(data))];kind='custom-continuous'
else:
data=mixture_data(4096);kind='four-gaussian-mixture'
split=int(.8*len(data));train=data[:split].to(args.device)
if args.method=='latent':
model,_,state,logs,stage_logs=train_latent(train,args)
elif args.method in ['diamond','meta']:
model,state,logs=train_posterior(args.method,train,args)
if args.finetune_steps:
drift,finetune_logs=train_reward_drift(model,train,args)
state['reward_drift']=drift.state_dict()
write_json(out/'finetune_losses.json',finetune_logs)
else:
model,state,logs,stage_logs=train_continuous(args.method,train,args)
state.update({'data_kind':kind,'heldout_data':data[split:]})
state.update({'format_version':1,'method':args.method,'width':args.width,'seed':args.seed})
torch.save(state,out/'checkpoint.pt')
write_json(out/'config.json',vars(args))
write_json(out/'losses.json',logs)
if stage_logs:write_json(out/'teacher_losses.json',stage_logs)
report={'method':args.method,'data_kind':state['data_kind'],
'train_loss_first_20_mean':sum(x['loss'] for x in logs[:20])/len(logs[:20]),
'train_loss_last_20_mean':sum(x['loss'] for x in logs[-20:])/len(logs[-20:]),
'training_steps':args.train_steps,'python':platform.python_version(),'torch':str(torch.__version__)}
if 'validation_ce' in state:report['heldout_diagonal_ce_at_half_time']=state['validation_ce']
if args.mode=='train-sample':
model.eval()
rows,metrics=generate(model,state,args)
(out/'samples.txt').write_text('\n'.join(rows)+'\n')
report.update(metrics)
report['elapsed_seconds']=time.perf_counter()-start
write_json(out/'report.json',report)
print(json.dumps({'method':args.method,'out':str(out),'loss':report['train_loss_last_20_mean'],
'seconds':report['elapsed_seconds']}),flush=True)
return report
if __name__=='__main__':main()