sra-trajectory-code / MID /main_nba_mid_graphv6_v2.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
26.9 kB
"""
main_nba_mid_graphv6_v2.py — MID on NBA with V6-style future interaction graph.
DESIGN PRINCIPLE
----------------
Only the denoising network changes from main_nba_mid.py:
- Past context encoder : NBAEncoder (GRU + SocialTransformer) — UNCHANGED
- Diffusion schedule : DiffusionTrajGraph (DDPM, 100 steps) — UNCHANGED
- Denoising network : GraphDenoiserNet (NEW)
TransformerConcatLinear backbone → eps_hat (noise prediction, same as baseline)
x_0 derived from eps for graph → geometry for neighbor selection
FutureInteractionGraphV6 (K=1) → refines eps_hat using inter-agent geometry
GRAPH INTEGRATION (same manner as V6)
--------------------------------------
Agent selection — RAG-style:
q_i = W_q([y0_emb_i, 0]) per-node query (no sigma: first pass only)
k_j = W_k([y0_emb_j, 0]) per-node key
semantic_score = (q_i · k_j) / √D_s
geo_bias = geo_mlp([mean_rel, std_rel, min_dist, heading_mean])
score = semantic + geo → top-N per target agent
Pairwise encoding — RelTrajEncoder:
edge_feat = RelTrajEncoder([rel_pos(2), heading_diff(1)]) over T steps
→ GNN message passing → gated residual refinement of x_0_hat
Usage:
python main_nba_mid_graphv6_v2.py --data_dir ../data/nba/original
"""
import os
import sys
import time
import logging
import argparse
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
from torch.utils.tensorboard import SummaryWriter # tbX-broken
from tqdm.auto import tqdm
# ---------------------------------------------------------------------------
# MoFlow graph module on sys.path
# ---------------------------------------------------------------------------
MOFLOW_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', 'MoFlow'))
sys.path.insert(0, MOFLOW_ROOT)
from models.graph_interaction_nba_v6 import FutureInteractionGraphV6
from models.context_encoder.mtr_encoder import SinusoidalPosEmb
# MID diffusion components (unchanged)
from models.diffusion import VarianceSchedule, TransformerConcatLinear
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
OBS_LEN = 10
PRED_LEN = 20
NUM_AGENTS = 11
TRAJ_SCALE = 5.0
TRAJ_MEAN = torch.FloatTensor([14.0, 7.5])
K_EVAL = 20
# ---------------------------------------------------------------------------
# Dataset (identical to main_nba_mid.py)
# ---------------------------------------------------------------------------
class NBADatasetMID(Dataset):
def __init__(self, data_dir: str, training: bool = True):
super().__init__()
fname = 'nba_train.npy' if training else 'nba_test.npy'
trajs = np.load(os.path.join(data_dir, fname)).astype(np.float32)
trajs /= (94.0 / 28.0)
trajs = torch.from_numpy(trajs).permute(0, 2, 1, 3) # (N, 11, 30, 2)
self.pre = trajs[:, :, :OBS_LEN, :]
self.fut = trajs[:, :, OBS_LEN:, :]
def __len__(self):
return len(self.pre)
def __getitem__(self, idx):
return self.pre[idx], self.fut[idx]
def nba_collate(batch):
pre = torch.stack([b[0] for b in batch]) # (B, 11, 10, 2)
fut = torch.stack([b[1] for b in batch]) # (B, 11, 20, 2)
return pre, fut
# ---------------------------------------------------------------------------
# Data pre-processing (identical to main_nba_mid.py)
# ---------------------------------------------------------------------------
def preprocess_batch(pre_motion, fut_motion, device):
"""
Returns:
past_6ch: [B*A, T_obs, 6] (÷ TRAJ_SCALE)
fut_rel: [B*A, T_fut, 2] (÷ TRAJ_SCALE, relative to last obs)
social_mask: [B*A, B*A]
last_obs: [B*A, 1, 2] un-scaled
"""
B, A = pre_motion.shape[:2]
traj_mean = TRAJ_MEAN.to(device)
pre = pre_motion.reshape(B * A, OBS_LEN, 2)
fut = fut_motion.reshape(B * A, PRED_LEN, 2)
last_obs = pre[:, -1:, :]
abs_xy = (pre - traj_mean) / TRAJ_SCALE
rel_xy = (pre - last_obs) / TRAJ_SCALE
vel_xy = torch.cat([rel_xy[:, 1:] - rel_xy[:, :-1],
torch.zeros_like(rel_xy[:, :1])], dim=1)
past_6ch = torch.cat([abs_xy, rel_xy, vel_xy], dim=-1)
fut_rel = (fut - last_obs) / TRAJ_SCALE
mask = torch.full((B * A, B * A), float('-inf'), device=device)
for i in range(B):
s, e = i * A, (i + 1) * A
mask[s:e, s:e] = 0.0
return past_6ch, fut_rel, mask, last_obs
# ---------------------------------------------------------------------------
# NBA Encoder (identical to main_nba_mid.py)
# ---------------------------------------------------------------------------
class _STEncoder(nn.Module):
def __init__(self, in_channels=6, hidden=256):
super().__init__()
self.conv = nn.Conv1d(in_channels, 32, kernel_size=3, padding=1)
self.relu = nn.ReLU()
self.gru = nn.GRU(32, hidden, num_layers=1, batch_first=True)
nn.init.kaiming_normal_(self.conv.weight)
nn.init.kaiming_normal_(self.gru.weight_ih_l0)
nn.init.kaiming_normal_(self.gru.weight_hh_l0)
nn.init.zeros_(self.conv.bias)
nn.init.zeros_(self.gru.bias_ih_l0)
nn.init.zeros_(self.gru.bias_hh_l0)
def forward(self, x):
h = self.relu(self.conv(x.transpose(1, 2)))
_, state = self.gru(h.transpose(1, 2))
return state.squeeze(0)
class _SocialTransformer(nn.Module):
def __init__(self, past_len=OBS_LEN, hidden=256):
super().__init__()
self.proj = nn.Linear(past_len * 6, hidden, bias=False)
layer = nn.TransformerEncoderLayer(
d_model=hidden, nhead=2,
dim_feedforward=hidden, batch_first=False)
self.encoder = nn.TransformerEncoder(layer, num_layers=2)
def forward(self, x_flat, mask):
h = self.proj(x_flat).unsqueeze(1)
h = h + self.encoder(h, mask)
return h.squeeze(1)
class NBAEncoder(nn.Module):
def __init__(self, encoder_dim=256, past_len=OBS_LEN):
super().__init__()
self.ego_encoder = _STEncoder(in_channels=6, hidden=256)
self.social_encoder = _SocialTransformer(past_len=past_len, hidden=256)
self.fusion = nn.Linear(512, encoder_dim)
def forward(self, past_6ch, social_mask):
ego = self.ego_encoder(past_6ch)
social = self.social_encoder(
past_6ch.reshape(past_6ch.size(0), -1), social_mask)
return self.fusion(torch.cat([ego, social], dim=-1))
# ---------------------------------------------------------------------------
# Graph-augmented denoising network
# ---------------------------------------------------------------------------
class GraphDenoiserNet(nn.Module):
"""TransformerConcatLinear backbone + V6-style future interaction graph.
Predicts epsilon (noise), same as original MID baseline.
Two-pass design (mirrors V6):
Pass 1 (no_grad):
backbone(x_t, beta, context) → eps_geom [N, T, 2]
x_0_geom = (x_t - c1*eps_geom) / c0 (detached)
Used ONLY for graph geometry — no gradient.
Pass 2 (grad-enabled):
backbone(x_t, beta, context) → eps_hat [N, T, 2] (grad-connected)
x_0_hat = (x_t - c1*eps_hat) / c0 (grad through eps_hat)
node_proj(x_0_hat) → y_emb [B, 1, A, D] (grad-connected)
FutureInteractionGraphV6 (K=1):
geometry : x_0_geom * TRAJ_SCALE + last_obs (detached from pass 1)
nodes : y_emb (grad-connected from pass 2)
→ RAG scoring (W_q/W_k + geo_bias) → top-N selection
→ RelTrajEncoder([rel_pos, heading_diff]) → GNN + gated residual
→ refined [B, 1, A, D]
refine_head(refined) → delta [N, T, 2]
output: eps_hat (pass 2) + delta
Gradient flow:
- Backbone receives gradients from eps_hat directly and via
x_0_hat → node_proj → graph → refine_head → delta.
- Edge geometry (positions) is detached → topology selection has no gradient.
"""
def __init__(self,
context_dim: int = 256,
tf_layer: int = 3,
T: int = PRED_LEN,
num_agents: int = NUM_AGENTS,
graph_hidden: int = 128,
top_n_neighbors: int = 5,
rel_traj_hidden: int = 32,
y0_score_dim: int = 32,
num_gnn_layers: int = 2,
graph_dropout: float = 0.1):
super().__init__()
self.T = T
self.A = num_agents
self.D = graph_hidden
# -- Backbone: standard MID denoising network ----------------------
self.backbone = TransformerConcatLinear(
point_dim = 2,
context_dim = context_dim,
tf_layer = tf_layer,
residual = False,
)
# -- Node projection: x_0_hat trajectory → graph embedding --------
self.node_proj = nn.Sequential(
nn.Linear(T * 2, graph_hidden),
nn.ReLU(inplace=True),
)
# -- Time embedding for GNN: beta (noise level) → [D] --------------
self.time_mlp = nn.Sequential(
SinusoidalPosEmb(graph_hidden),
nn.Linear(graph_hidden, graph_hidden),
nn.ReLU(),
nn.Linear(graph_hidden, graph_hidden),
)
# -- V6-style future interaction graph (K=1) -----------------------
self.future_graph = FutureInteractionGraphV6(
embed_dim = graph_hidden,
future_steps = T,
num_agents = num_agents,
num_heads = 4,
dropout = graph_dropout,
num_gnn_layers = num_gnn_layers,
time_dim = graph_hidden,
top_n_neighbors = top_n_neighbors,
rel_traj_hidden = rel_traj_hidden,
y0_score_dim = y0_score_dim,
)
# -- Refinement head: embedding → trajectory correction ------------
self.refine_head = nn.Linear(graph_hidden, T * 2)
# Zero-init: graph starts as an identity (no correction), training adds it
nn.init.zeros_(self.refine_head.weight)
nn.init.zeros_(self.refine_head.bias)
def forward(self,
x_t: torch.Tensor, # [N=B*A, T, 2] noisy trajectory
beta: torch.Tensor, # [N] noise level (same per scene)
context: torch.Tensor, # [N, context_dim]
B: int,
A: int,
last_obs: torch.Tensor, # [N, 1, 2] last observed position (un-scaled)
tau: torch.Tensor, # [B] normalised timestep ∈ [0, 1]
alpha_bar: torch.Tensor, # [N] cumulative alpha for x_0 recovery
) -> torch.Tensor: # [N, T, 2] predicted noise (epsilon)
N, T, _ = x_t.shape
D = self.D
c0 = alpha_bar.sqrt().view(N, 1, 1) # [N, 1, 1]
c1 = (1 - alpha_bar).sqrt().view(N, 1, 1)
# ---- Pass 1 (no_grad): predict eps → derive x_0 for geometry -----
with torch.no_grad():
eps_geom = self.backbone(x_t, beta=beta, context=context) # [N, T, 2]
x_0_geom = (x_t - c1 * eps_geom) / c0 # [N, T, 2]
# Absolute positions (detached — geometry has no gradient)
x_0_abs = x_0_geom * TRAJ_SCALE + last_obs # [N, T, 2]
y_abs = x_0_abs.view(B, A, T, 2).unsqueeze(1) # [B, 1, A, T, 2]
# ---- Pass 2 (grad-enabled): predict eps + derive x_0 for nodes ---
eps_hat = self.backbone(x_t, beta=beta, context=context) # [N, T, 2]
x_0_hat = (x_t - c1 * eps_hat) / c0 # grad through eps_hat
# ---- Node embeddings from derived x_0 (grad-connected) -----------
y_emb = self.node_proj(x_0_hat.view(N, T * 2)) # [N, D]
y_emb = y_emb.view(B, 1, A, D) # [B, 1, A, D]
# ---- 4. Time embedding (per scene, from per-agent beta) -----------
beta_scene = beta.view(B, A)[:, 0] # [B]
t_emb = self.time_mlp(beta_scene) # [B, D]
# ---- 5. V6 future interaction graph (K=1, sigma=None) ------------
refined = self.future_graph(
y_emb, y_abs, t_emb, tau, sigma_agent=None,
) # [B, 1, A, D]
# ---- 6. Decode and add residual correction ------------------------
delta = self.refine_head(
refined.squeeze(1).reshape(N, D)
).view(N, T, 2) # [N, T, 2]
return eps_hat + delta
# ---------------------------------------------------------------------------
# Graph-aware diffusion wrapper
# ---------------------------------------------------------------------------
class DiffusionTrajGraph(nn.Module):
"""DDPM wrapper for GraphDenoiserNet.
Key differences from DiffusionTraj:
- Predicts epsilon (noise), same as original MID
- Derives x_0 from epsilon internally for graph geometry
- Samples one timestep per scene (not per agent)
→ all 11 agents in a scene share the same noise level,
giving the graph a consistent geometric signal
- Passes B, A, last_obs, tau, alpha_bar to the net at every step
"""
def __init__(self, net: GraphDenoiserNet, var_sched: VarianceSchedule):
super().__init__()
self.net = net
self.var_sched = var_sched
def get_loss(self,
x_0: torch.Tensor, # [B*A, T, 2] ground-truth relative future
context: torch.Tensor, # [B*A, enc_dim]
B: int,
A: int,
last_obs: torch.Tensor, # [B*A, 1, 2] un-scaled last obs
) -> torch.Tensor:
device = x_0.device
N = B * A
# One timestep per scene, repeated for each of its A agents
t_scene = self.var_sched.uniform_sample_t(B) # list of B ints ∈ [1, T]
t = [ts for ts in t_scene for _ in range(A)] # B*A ints
alpha_bar = self.var_sched.alpha_bars[t].to(device) # [N]
beta = self.var_sched.betas[t].to(device) # [N]
c0 = alpha_bar.sqrt().view(N, 1, 1) # [N, 1, 1]
c1 = (1 - alpha_bar).sqrt().view(N, 1, 1)
e_rand = torch.randn_like(x_0)
x_t = c0 * x_0 + c1 * e_rand # [N, T, 2]
tau = torch.tensor(
[ts / self.var_sched.num_steps for ts in t_scene],
dtype=torch.float32, device=device,
) # [B] ∈ [0, 1]
eps_pred = self.net(x_t, beta, context, B, A, last_obs, tau, alpha_bar)
return F.mse_loss(eps_pred, e_rand)
@torch.no_grad()
def sample(self,
num_points: int,
context: torch.Tensor, # [B*A, enc_dim]
B: int,
A: int,
last_obs: torch.Tensor, # [B*A, 1, 2]
sample: int = K_EVAL,
bestof: bool = True,
sampling: str = 'ddim',
step: int = 10,
) -> torch.Tensor: # [K, B*A, T, 2]
device = context.device
N = B * A
stride = self.var_sched.num_steps // step
traj_list = []
for _ in range(sample):
x_T = (torch.randn(N, num_points, 2, device=device)
if bestof else
torch.zeros(N, num_points, 2, device=device))
x_t = x_T
for t in range(self.var_sched.num_steps, 0, -stride):
alpha_bar = self.var_sched.alpha_bars[t].to(device)
alpha_bar_prev = self.var_sched.alpha_bars[t - stride].to(device)
beta_t = self.var_sched.betas[t].to(device)
beta_batch = beta_t.expand(N)
alpha_bar_vec = alpha_bar.expand(N)
tau = torch.full((B,), t / self.var_sched.num_steps,
dtype=torch.float32, device=device)
e_theta = self.net(x_t, beta_batch, context, B, A, last_obs,
tau, alpha_bar_vec)
if sampling == 'ddim':
x0_t = (x_t - (1 - alpha_bar).sqrt() * e_theta) / alpha_bar.sqrt()
x_t = (alpha_bar_prev.sqrt() * x0_t
+ (1 - alpha_bar_prev).sqrt() * e_theta)
elif sampling == 'ddpm':
sigma = self.var_sched.get_sigmas(t, flexibility=0.0)
z = torch.randn_like(x_t) if t > stride else torch.zeros_like(x_t)
c0 = 1.0 / self.var_sched.alphas[t].sqrt()
c1 = (1 - self.var_sched.alphas[t]) / (1 - alpha_bar).sqrt()
x_t = c0 * (x_t - c1 * e_theta) + sigma * z
else:
raise ValueError(f"Unknown sampling: {sampling}")
traj_list.append(x_t) # [N, T, 2]
return torch.stack(traj_list) # [K, N, T, 2]
# ---------------------------------------------------------------------------
# Trainer
# ---------------------------------------------------------------------------
class Trainer:
def __init__(self, args):
self.args = args
self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
self._build_dirs()
self._build_data()
self._build_model()
self._build_optimizer()
def _build_dirs(self):
self.exp_dir = os.path.join('experiments', self.args.exp_name)
os.makedirs(self.exp_dir, exist_ok=True)
self.tb_log = SummaryWriter(log_dir=self.exp_dir)
log_path = os.path.join(
self.exp_dir,
'nba_{}.log'.format(time.strftime('%Y-%m-%d-%H-%M')))
self.log = logging.getLogger(self.args.exp_name)
self.log.setLevel(logging.INFO)
self.log.addHandler(logging.FileHandler(log_path))
self.log.addHandler(logging.StreamHandler(sys.stdout))
self.log.info(f"Args: {self.args}")
def _build_data(self):
train_dset = NBADatasetMID(self.args.data_dir, training=True)
test_dset = NBADatasetMID(self.args.data_dir, training=False)
self.train_loader = DataLoader(
train_dset, batch_size=self.args.batch_size,
shuffle=True, num_workers=4,
collate_fn=nba_collate, pin_memory=True)
self.test_loader = DataLoader(
test_dset, batch_size=self.args.eval_batch_size,
shuffle=False, num_workers=4,
collate_fn=nba_collate, pin_memory=True)
self.log.info(
f"Train: {len(train_dset)} scenes Test: {len(test_dset)} scenes")
def _build_model(self):
# ---- Past context encoder (unchanged from main_nba_mid.py) --------
self.encoder = NBAEncoder(
encoder_dim = self.args.encoder_dim,
past_len = OBS_LEN,
).to(self.device)
# ---- Denoising network with V6 graph interaction ------------------
net = GraphDenoiserNet(
context_dim = self.args.encoder_dim,
tf_layer = self.args.tf_layer,
T = PRED_LEN,
num_agents = NUM_AGENTS,
graph_hidden = self.args.graph_hidden,
top_n_neighbors = self.args.top_n_neighbors,
rel_traj_hidden = self.args.rel_traj_hidden,
y0_score_dim = self.args.y0_score_dim,
num_gnn_layers = self.args.graph_gnn_layers,
graph_dropout = self.args.graph_dropout,
)
self.diffusion = DiffusionTrajGraph(
net = net,
var_sched = VarianceSchedule(
num_steps = 100,
beta_T = 5e-2,
mode = 'linear',
),
).to(self.device)
n_enc = sum(p.numel() for p in self.encoder.parameters())
n_back = sum(p.numel() for p in net.backbone.parameters())
n_grph = sum(p.numel() for p in net.future_graph.parameters())
n_diff = sum(p.numel() for p in self.diffusion.parameters())
self.log.info(f"Encoder params: {n_enc:,}")
self.log.info(f"Backbone params: {n_back:,}")
self.log.info(f"Graph params: {n_grph:,}")
self.log.info(f"Total diffusion params:{n_diff:,}")
def _build_optimizer(self):
params = (list(self.encoder.parameters()) +
list(self.diffusion.parameters()))
self.optimizer = torch.optim.Adam(params, lr=self.args.lr)
self.scheduler = torch.optim.lr_scheduler.ExponentialLR(
self.optimizer, gamma=0.98)
# ------------------------------------------------------------------
def train(self):
for epoch in range(1, self.args.epochs + 1):
self.encoder.train()
self.diffusion.train()
total_loss, count = 0.0, 0
pbar = tqdm(self.train_loader, ncols=90)
for pre, fut in pbar:
pre = pre.to(self.device)
fut = fut.to(self.device)
B = pre.size(0)
past_6ch, fut_rel, mask, last_obs = preprocess_batch(
pre, fut, self.device)
context = self.encoder(past_6ch, mask)
loss = self.diffusion.get_loss(
fut_rel, context, B=B, A=NUM_AGENTS, last_obs=last_obs)
self.optimizer.zero_grad()
loss.backward()
nn.utils.clip_grad_norm_(
list(self.encoder.parameters()) +
list(self.diffusion.parameters()), 1.0)
self.optimizer.step()
total_loss += loss.item()
count += 1
pbar.set_description(
f"Epoch {epoch} loss={total_loss/count:.4f}")
self.scheduler.step()
avg_loss = total_loss / count
self.tb_log.add_scalar('loss/train', avg_loss, epoch)
self.log.info(f"Epoch {epoch} train_loss={avg_loss:.4f}")
if epoch % self.args.eval_every == 0:
ade, fde = self.evaluate()
self.tb_log.add_scalar('metric/ADE', ade, epoch)
self.tb_log.add_scalar('metric/FDE', fde, epoch)
self.log.info(f"Epoch {epoch} ADE={ade:.4f} FDE={fde:.4f}")
torch.save({
'encoder': self.encoder.state_dict(),
'diffusion': self.diffusion.state_dict(),
'epoch': epoch,
}, os.path.join(self.exp_dir, f'nba_epoch{epoch:04d}.pt'))
@torch.no_grad()
def evaluate(self):
self.encoder.eval()
self.diffusion.eval()
ade_sum, fde_sum, n_agents = 0.0, 0.0, 0
for pre, fut in tqdm(self.test_loader, ncols=90, desc='Eval'):
pre = pre.to(self.device)
fut = fut.to(self.device)
B = pre.size(0)
past_6ch, _, mask, last_obs = preprocess_batch(
pre, fut, self.device)
context = self.encoder(past_6ch, mask) # [B*11, enc_dim]
# pred_rel: [K, B*11, T, 2] (in ÷TRAJ_SCALE space)
pred_rel = self.diffusion.sample(
num_points = PRED_LEN,
context = context,
B = B,
A = NUM_AGENTS,
last_obs = last_obs,
sample = K_EVAL,
bestof = True,
sampling = self.args.sampling,
step = self.args.sampling_step,
)
# absolute positions
pred_abs = pred_rel * TRAJ_SCALE + last_obs.unsqueeze(0) # [K, B*11, T, 2]
fut_abs = fut.reshape(B * NUM_AGENTS, PRED_LEN, 2) # [B*11, T, 2]
dist = (pred_abs - fut_abs.unsqueeze(0)).norm(dim=-1) # [K, B*11, T]
# minADE: pick best single trajectory (mean over T, then min over K)
ade_per_mode = dist.mean(dim=-1) # [K, B*11]
ade_sum += ade_per_mode.min(dim=0).values.sum().item()
# minFDE: pick best mode at last frame
fde_sum += dist[:, :, -1].min(dim=0).values.sum().item()
n_agents += B * NUM_AGENTS
ade = ade_sum / n_agents
fde = fde_sum / n_agents
return ade, fde
# ---------------------------------------------------------------------------
# Argument parsing
# ---------------------------------------------------------------------------
def parse_args():
p = argparse.ArgumentParser()
# Data
p.add_argument('--data_dir', type=str, default='../data/nba/original')
# Experiment
p.add_argument('--exp_name', type=str, default='mid_nba_graphv6v2')
# Training
p.add_argument('--epochs', type=int, default=100)
p.add_argument('--batch_size', type=int, default=32)
p.add_argument('--eval_batch_size', type=int, default=64)
p.add_argument('--lr', type=float, default=1e-3)
p.add_argument('--eval_every', type=int, default=5)
# MID backbone
p.add_argument('--encoder_dim', type=int, default=256)
p.add_argument('--tf_layer', type=int, default=3)
# Graph
p.add_argument('--graph_hidden', type=int, default=128)
p.add_argument('--graph_gnn_layers', type=int, default=2)
p.add_argument('--graph_dropout', type=float, default=0.1)
p.add_argument('--top_n_neighbors', type=int, default=5)
p.add_argument('--rel_traj_hidden', type=int, default=32)
p.add_argument('--y0_score_dim', type=int, default=32)
# Sampling
p.add_argument('--sampling', type=str, default='ddim',
choices=['ddpm', 'ddim'])
p.add_argument('--sampling_step', type=int, default=10)
return p.parse_args()
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
if __name__ == '__main__':
args = parse_args()
trainer = Trainer(args)
trainer.train()