Download lecture_6/run.py from ChatterjeeLab/CIS6270: direct link, hf CLI and curl.
- Browser
- Download file 14.1 kB
-
https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_6/run.py
- Command line
-
hf download hf://ChatterjeeLab/CIS6270/lecture_6/run.py
-
curl -L -o run.py https://huggingface.co/ChatterjeeLab/CIS6270/resolve/main/lecture_6/run.py
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() | |