Spaces:
Running on Zero
Running on Zero
File size: 2,906 Bytes
2680bd5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 | from dataclasses import dataclass
from typing import List, Literal
@dataclass
class DiffusionConfig:
noise_schedule: str
diffusion_steps: int
sigma_small: bool
repr: Literal['joint_pos', 'joint_rot', 'joint_pos_w_scalar_rot', 'joint_pos_w_axisangle_rot']
contact_loss: bool
lambda_contact: float | None
lambda_contact_predict: float | None
lambda_rcxyz: float
lambda_repr: float
lambda_vel: float
lambda_acce: float
lambda_fc: float
lambda_w_ig: float
@dataclass
class ModelConfig:
latent_dim: int
layers: int
num_heads: int
ff_size: int
dropout: float
activation: str
cond_mode: str
diffusion: DiffusionConfig
contact_prediction: bool
repr: Literal['joint_pos', 'joint_rot', 'joint_pos_w_scalar_rot', 'joint_pos_w_axisangle_rot']
@dataclass
class ActionConditionModelConfig(ModelConfig):
cond_mask_prob: float
num_actions: int
@dataclass
class TextConditionModelConfig(ModelConfig):
arch: str
text_model: str
max_text_length: int | None
cond_mask_prob: float
treble_mask_prob: float
@dataclass
class DataConfig:
ratio: float
fixed_length: int
max_length: int
min_length: int
normalize: bool
difference: bool
repr: Literal['joint_pos', 'joint_rot', 'joint_pos_w_scalar_rot', 'joint_pos_w_axisangle_rot']
contact_label: bool
data_dir: str
use_plain: bool | None
data_file_name: str | None
@dataclass
class DataLoaderConfig:
batch_size: int
num_workers: int
shuffle: bool
@dataclass
class OptimizerConfig:
lr: float
weight_decay: float
@dataclass
class SampleConfig:
guidance_param: float
@dataclass
class VisualizationConfig:
denoising_steps: List[int]
samples_count: int
@dataclass
class ValidationConfig:
val_interval: int
dataloader: DataLoaderConfig
@dataclass
class EvaluationConfig:
dataloader: DataLoaderConfig
eval_interval: int
num_samples_on_train: int
num_samples_on_val: int
num_samples_per_condition: int
@dataclass
class TrainingConfig:
save_dir: str
overwrite: bool
train_platform_type: str
log_interval: int
save_interval: int
num_steps: int
resume_checkpoint: str
eval_during_training: bool
eval_cfg: EvaluationConfig | None
val_during_training: bool
val_cfg: ValidationConfig | None
optimizer: OptimizerConfig
sample: SampleConfig
dataloader: DataLoaderConfig
viz_during_training: bool
viz_cfg: VisualizationConfig | None
@dataclass
class Config:
seed: int
model: ModelConfig | ActionConditionModelConfig | TextConditionModelConfig
data: DataConfig
train: TrainingConfig
@dataclass
class GenerateConfig:
model_path: str
output_dir: str
num_samples: int
sample: SampleConfig
action_name: str | None
text_prompt: str | None
motion_length: int |