ScanDL2 / scandl_module /scripts /sp_basic_utils.py
Di0nigi's picture
First commit
95456ed verified
Raw
History Blame
3 kB
import argparse
import json
import sys
from ScanDL2.scandl_module.original_scandl import sp_gaussian_diffusion as gd
from ScanDL2.scandl_module.original_scandl.sp_gaussian_diffusion import (
SpacedDiffusion,
space_timesteps,
)
from ScanDL2.scandl_module.original_scandl.sp_transformer_model import TransformerNetModel
sys.path.append("./")
sys.path.append("../")
def load_defaults_config(config_path: str):
"""
Load defaults for training args.
"""
with open(config_path, "r") as f:
return json.load(f)
def create_model_and_diffusion(
hidden_t_dim,
hidden_dim,
vocab_size,
config_name,
use_plm_init,
dropout,
num_transformer_layers,
num_transformer_heads,
mask_padding,
diffusion_steps,
noise_schedule,
learn_sigma,
timestep_respacing,
predict_xstart,
rescale_timesteps,
sigma_small,
rescale_learned_sigmas,
use_kl,
one_noise_step,
nll_in_loss,
notes,
**kwargs,
):
model = TransformerNetModel(
input_dims=hidden_dim,
output_dims=(hidden_dim if not learn_sigma else hidden_dim * 2),
hidden_t_dim=hidden_t_dim,
num_transformer_layers=num_transformer_layers,
num_transformer_heads=num_transformer_heads,
one_noise_step=one_noise_step,
mask_padding=mask_padding,
dropout=dropout,
config_name=config_name,
vocab_size=vocab_size,
init_pretrained=use_plm_init,
)
betas = gd.get_named_beta_schedule(noise_schedule, diffusion_steps)
if not timestep_respacing:
timestep_respacing = [diffusion_steps]
diffusion = SpacedDiffusion(
use_timesteps=space_timesteps(diffusion_steps, timestep_respacing),
betas=betas,
rescale_timesteps=rescale_timesteps,
predict_xstart=predict_xstart,
learn_sigmas=learn_sigma,
sigma_small=sigma_small,
use_kl=use_kl,
one_noise_step=one_noise_step,
nll_in_loss=nll_in_loss,
mask_padding=mask_padding,
rescale_learned_sigmas=rescale_learned_sigmas,
)
return model, diffusion
def add_dict_to_argparser(parser, default_dict):
for k, v in default_dict.items():
v_type = type(v)
if v is None:
v_type = str
elif isinstance(v, bool):
v_type = str2bool
parser.add_argument(f"--{k}", default=v, type=v_type)
def args_to_dict(args, keys):
return {k: getattr(args, k) for k in keys}
def str2bool(v):
"""
https://stackoverflow.com/questions/15008758/parsing-boolean-values-with-argparse
"""
if isinstance(v, bool):
return v
if v.lower() in ("yes", "true", "t", "y", "1"):
return True
elif v.lower() in ("no", "false", "f", "n", "0"):
return False
else:
raise argparse.ArgumentTypeError("boolean value expected")