Skip to content

Commit 74ef1bd

Browse files
authored
Move date preprocessing into Energon for RoboBrain-X0 training. (flagos-ai#1069)
### PR Category Train ### PR Types Improvements ### PR Description Move date preprocessing into Energon for RoboBrain-X0 training.
1 parent d5bdeca commit 74ef1bd

7 files changed

Lines changed: 216 additions & 154 deletions

File tree

examples/robobrain_x0_5/conf/train/libero_qwengroot.yaml

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,9 @@ output_directory: outputs/libero_qwengroot/checkpoints/ckpt_out
1414
log_freq: 10
1515
train_steps: 100
1616

17+
shuffle_buffer_size: null
18+
shuffle_over_epochs_multiplier: -1
19+
1720
framework:
1821
name: QwenGR00T
1922
qwenvl:

flagscale/models/robobrain_x/qwen_groot.py

Lines changed: 5 additions & 70 deletions
Original file line numberDiff line numberDiff line change
@@ -65,19 +65,7 @@ def __init__(self, config: dict | None = None, **kwargs) -> None:
6565
self.past_action_window_size = config.framework.action_model.past_action_window_size
6666
self.chunk_len = self.past_action_window_size + 1 + self.future_action_window_size
6767

68-
def forward(self, examples: list[dict] | None = None, **kwargs) -> tuple:
69-
batch_images = [example["image"] for example in examples] # [B,[PLT]]
70-
instructions = [example["lang"] for example in examples] # [B, str]
71-
actions = [example["action"] for example in examples] # label [B, len, 7]
72-
73-
state = (
74-
[example["state"] for example in examples] if "state" in examples[0] else None
75-
) # [B, 1, state_dim]
76-
77-
# Step 1: QWenVL input format
78-
qwen_inputs = self.qwen_vl_interface.build_qwenvl_inputs(
79-
images=batch_images, instructions=instructions
80-
)
68+
def forward(self, qwen_inputs, state, actions, **kwargs) -> tuple:
8169
with torch.autocast("cuda", dtype=torch.bfloat16):
8270
qwenvl_outputs = self.qwen_vl_interface(
8371
**qwen_inputs, output_attentions=False, output_hidden_states=True, return_dict=True
@@ -88,14 +76,7 @@ def forward(self, examples: list[dict] | None = None, **kwargs) -> tuple:
8876
# Step 4: Action Expert Forward and Loss
8977
with torch.autocast("cuda", dtype=torch.float32):
9078
# [B, T_full, action_dim]
91-
if isinstance(actions[0], torch.Tensor):
92-
actions = torch.stack(actions, dim=0).to(
93-
device=last_hidden.device, dtype=last_hidden.dtype
94-
)
95-
else:
96-
actions = torch.tensor(
97-
np.array(actions), device=last_hidden.device, dtype=last_hidden.dtype
98-
)
79+
actions = actions.to(device=last_hidden.device, dtype=last_hidden.dtype)
9980
actions_target = actions[
10081
:, -(self.future_action_window_size + 1) :, :
10182
] # (B, chunk_len, action_dim)
@@ -109,16 +90,8 @@ def forward(self, examples: list[dict] | None = None, **kwargs) -> tuple:
10990
last_hidden_repeated = last_hidden.repeat(repeated_diffusion_steps, 1, 1)
11091

11192
state_repeated = None
112-
if state is not None:
113-
if isinstance(state[0], torch.Tensor):
114-
state = torch.stack(state, dim=0).to(
115-
device=last_hidden.device, dtype=last_hidden.dtype
116-
)
117-
else:
118-
state = torch.tensor(
119-
np.array(state), device=last_hidden.device, dtype=last_hidden.dtype
120-
)
121-
state_repeated = state.repeat(repeated_diffusion_steps, 1, 1)
93+
state = state.to(device=last_hidden.device, dtype=last_hidden.dtype)
94+
state_repeated = state.repeat(repeated_diffusion_steps, 1, 1)
12295

12396
action_loss = self.action_model(
12497
last_hidden_repeated, actions_target_repeated, state_repeated
@@ -297,38 +270,6 @@ def dryrun_with_random_sample(cfg):
297270
model.save_pretrained()
298271

299272

300-
def dryrun_with_dataloader(cfg):
301-
model: Qwen_GR00T = Qwen_GR00T(cfg)
302-
303-
from megatron.energon import WorkerConfig, get_loader, get_train_dataset
304-
from tools.datasets.vla.data.dataset_helpers_np_pil import TaskEncoder
305-
306-
ds = get_train_dataset(
307-
cfg.datasets.data_path,
308-
batch_size=1,
309-
shuffle_buffer_size=100,
310-
max_samples_per_sequence=100,
311-
worker_config=WorkerConfig.default_worker_config(num_workers=1, data_parallel_group=None),
312-
task_encoder=TaskEncoder(cfg.datasets.task_encoder),
313-
repeat=True,
314-
)
315-
vla_train_dataloader = get_loader(ds)
316-
data_iter = iter(vla_train_dataloader)
317-
batch = next(data_iter)
318-
batch = get_batch(batch)
319-
320-
# try get model
321-
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
322-
model = model.to(device)
323-
model(batch)
324-
forward_output = model(batch)
325-
action_loss = forward_output["action_loss"]
326-
print(f"Action Loss: {action_loss.item()}")
327-
328-
action = model.predict_action(batch_images=[batch[0]["image"]], instructions=[batch[0]["lang"]])
329-
print(f"Action inference: {action['normalized_actions'].shape}")
330-
331-
332273
if __name__ == "__main__":
333274
import argparse
334275

@@ -339,13 +280,7 @@ def dryrun_with_dataloader(cfg):
339280
default="./examples/robobrain_x0_5/conf/train/libero_qwengroot.yaml",
340281
help="Path to YAML config",
341282
)
342-
parser.add_argument("--dryrun-dataloader", action="store_true")
343-
parser.add_argument("--dryrun-random", action="store_true")
344283

345284
args, clipargs = parser.parse_known_args()
346285
cfg = OmegaConf.load(args.config_yaml)
347-
348-
if args.dryrun_dataloader:
349-
dryrun_with_dataloader(cfg)
350-
if args.dryrun_random:
351-
dryrun_with_random_sample(cfg)
286+
dryrun_with_random_sample(cfg)

flagscale/runner/runner_train.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -464,6 +464,7 @@ def run(
464464
monitor=False,
465465
interval=10,
466466
enable_monitoring=None,
467+
**kwargs,
467468
):
468469
# Read from config if not explicitly provided
469470
if enable_monitoring is None:

flagscale/train/megatron/train_robobrain_x0.5_qwengroot.py

Lines changed: 34 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
import pathlib
1212
import platform
1313
import random
14+
import time
1415

1516
import epath
1617
import numpy as np
@@ -22,10 +23,10 @@
2223

2324
import wandb
2425
from megatron.energon import WorkerConfig, get_loader, get_train_dataset
25-
from tools.datasets.vla.data.dataset_helpers_np_pil import TaskEncoder
26+
from tools.datasets.vla.data.dataset_helpers_preprocess import TaskEncoder
2627

2728
from flagscale.logger import logger
28-
from flagscale.models.robobrain_x.qwen_groot import Qwen_GR00T, get_batch
29+
from flagscale.models.robobrain_x.qwen_groot import Qwen_GR00T
2930

3031
# Sane Defaults
3132
os.environ["TOKENIZERS_PARALLELISM"] = "false"
@@ -94,15 +95,20 @@ def setup_optimizer_and_scheduler(
9495

9596

9697
def init_ddp(seed):
97-
torch.manual_seed(seed)
98-
np.random.seed(seed)
98+
os.environ["PYTHONHASHSEED"] = str(seed)
99+
os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
99100
random.seed(seed)
100-
local_rank = int(os.environ["LOCAL_RANK"])
101-
torch.cuda.set_device(local_rank)
102-
torch.distributed.init_process_group(backend='nccl', init_method='env://')
101+
np.random.seed(seed)
102+
torch.manual_seed(seed)
103+
torch.cuda.manual_seed_all(seed)
103104
torch.backends.cudnn.enabled = True
104-
torch.backends.cudnn.benchmark = True
105+
torch.backends.cudnn.benchmark = False
105106
torch.backends.cudnn.deterministic = True
107+
torch.use_deterministic_algorithms(True)
108+
109+
local_rank = int(os.environ.get("LOCAL_RANK", 0))
110+
torch.cuda.set_device(local_rank)
111+
torch.distributed.init_process_group(backend="nccl", init_method="env://")
106112
return local_rank
107113

108114

@@ -126,16 +132,19 @@ def init_wandb(config, *, resuming: bool, log_code: bool = False, enabled: bool
126132

127133

128134
def main(cfg) -> None:
135+
local_rank = init_ddp(cfg.seed)
136+
129137
# build model
130138
vla = Qwen_GR00T(cfg)
131139
# prepare data
132140
ds = get_train_dataset(
133141
cfg.datasets.data_path,
134142
batch_size=cfg.batch_size,
135-
shuffle_buffer_size=100,
143+
shuffle_buffer_size=cfg.shuffle_buffer_size,
136144
max_samples_per_sequence=100,
145+
shuffle_over_epochs_multiplier=cfg.shuffle_over_epochs_multiplier,
137146
worker_config=WorkerConfig.default_worker_config(num_workers=1, data_parallel_group=None),
138-
task_encoder=TaskEncoder(cfg.datasets.task_encoder),
147+
task_encoder=TaskEncoder(cfg),
139148
repeat=True,
140149
)
141150
vla_train_dataloader = get_loader(ds)
@@ -145,13 +154,9 @@ def main(cfg) -> None:
145154
# set optimizer and scheduler
146155
optimizer, lr_scheduler = setup_optimizer_and_scheduler(model=vla, cfg=cfg)
147156
# Run VLA Training
148-
local_rank = init_ddp(cfg.seed)
157+
149158
if dist.get_rank() == 0 and local_rank == 0:
150159
logger.info(f"Running on: {platform.node()}")
151-
if cfg.batch_size % torch.cuda.device_count() != 0:
152-
raise ValueError(
153-
f"Batch size {cfg.batch_size} must be divisible by the number of devices {torch.cuda.device_count()}."
154-
)
155160
resuming = cfg.resume
156161
init_wandb(cfg, resuming=resuming, enabled=cfg.wandb_enabled)
157162

@@ -160,18 +165,28 @@ def main(cfg) -> None:
160165

161166
step = 0
162167
done = False
168+
169+
t_start = time.time()
163170
while not done:
164171
batch = next(data_iter)
165-
batch = get_batch(batch)
166-
output_dict = vla.forward(batch)
172+
173+
qwen_inputs, state, actions = batch.get("qwen_inputs"), batch.get("state"), batch.get("actions")
174+
if not qwen_inputs or not actions:
175+
continue
176+
for i in qwen_inputs:
177+
qwen_inputs[i] = qwen_inputs[i].to(device=vla.device)
178+
output_dict = vla.forward(qwen_inputs=qwen_inputs, state=state, actions=actions)
179+
167180
action_loss = output_dict["action_loss"]
168181
action_loss.backward()
169182
optimizer.step()
170183
lr_scheduler.step()
171184
optimizer.zero_grad()
172185

173-
if step % cfg.log_freq == 0:
174-
logger.info(f"step: {step} loss: {action_loss.item():.3f}")
186+
if step % cfg.log_freq == 0 and dist.get_rank() == 0 and local_rank == 0:
187+
logger.info(f"step {step} loss: {action_loss.item()}")
188+
logger.info(f"step {step}: {(time.time() - t_start) / cfg.log_freq:.3f}s/iter")
189+
t_start = time.time()
175190
step += 1
176191
if step >= cfg.train_steps:
177192
done = True

flagscale/train/train_pi.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -76,7 +76,7 @@ def set_seed(seed: int):
7676

7777

7878
def init_ddp():
79-
local_rank = int(os.environ["LOCAL_RANK"])
79+
local_rank = int(os.environ.get("LOCAL_RANK", 0))
8080
torch.cuda.set_device(local_rank)
8181
torch.distributed.init_process_group(backend="nccl", init_method="env://")
8282

tools/datasets/vla/data/dataset_helpers_np_pil.py

Lines changed: 0 additions & 64 deletions
This file was deleted.

0 commit comments

Comments
 (0)