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