Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions torchtitan/experiments/worldmodel/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,8 @@ class _DiffusionConfig:
train_skip: int
val_skip: int
nan_engaged_plans: bool
image_size: tuple[int, int] = (128, 256)
action_targets: bool = False

def skip(self, val: bool) -> int:
return self.val_skip if val else self.train_skip
Expand Down Expand Up @@ -75,6 +77,13 @@ def _mock_batch(self) -> tuple[dict[str, np.ndarray], dict[str, np.ndarray]]:
}
)
targets = {"plan": np.random.randn(batch, PLAN_SIZE).astype(np.float32)}
if cfg.action_targets:
from xx.training.path.dataloader import compute_action_target_from_plan

inputs["action_t"] = np.random.uniform(0, 1, (batch, 2)).astype(np.float32)
targets["action"] = compute_action_target_from_plan(
targets["plan"][:, :495].reshape(-1, 33, 15), inputs["action_t"]
)
return inputs, targets


Expand Down Expand Up @@ -106,6 +115,7 @@ class Config(BaseDataLoader.Config):
limit: int | None
mock_data: bool
mock_segment_batch_size: int
action_targets: bool = False

def __post_init__(self) -> None:
total_frames = self.context_size_frames + self.future_size_frames
Expand Down Expand Up @@ -226,6 +236,8 @@ def _build_dataset(config: Config, *, val: bool, global_rank: int, global_world_
train_skip=config.train_skip,
val_skip=config.val_skip,
nan_engaged_plans=config.nan_engaged_plans,
image_size=config.image_size,
action_targets=config.action_targets,
),
val=val,
local_rank=int(os.environ.get("LOCAL_RANK", "0")),
Expand Down
53 changes: 35 additions & 18 deletions torchtitan/experiments/worldmodel/loss.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,9 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from __future__ import annotations

import math
Expand Down Expand Up @@ -41,30 +47,39 @@ def compute_worldmodel_losses(
targets: dict[str, torch.Tensor],
*,
plan_loss_weight: float,
action_loss_weight: float = 1.0,
) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
loss: torch.Tensor | None = None
terms: dict[str, torch.Tensor] = {}

if "sample" in outputs:
pred = outputs["sample"]
target = targets["v"].to(device=pred.device, dtype=pred.dtype)
mask = targets["mask"].to(device=pred.device).flatten(1).float()
for output, target_key, mask_key, weight, term in (
("sample", "v", "mask", 1.0, "diffusion_loss"),
("plan_v", "plan_v", "plan_mask", plan_loss_weight, "plan_diffusion_loss"),
):
if output not in outputs:
continue
pred = outputs[output] if output == "sample" else outputs[output][:, -1]
target = targets[target_key].to(device=pred.device, dtype=pred.dtype)
mask = targets[mask_key].to(device=pred.device).flatten(1).float()
mse = F.mse_loss(pred.float(), target.float(), reduction="none").flatten(1)
diffusion_loss = (mse * mask).sum(dim=1) / mask.sum(dim=1).clamp_min(1.0)
loss = diffusion_loss
terms["diffusion_loss"] = diffusion_loss.detach()

if "plan" in outputs and "plan" in targets:
pred = outputs["plan"]
target = targets["plan"].to(device=pred.device, dtype=pred.dtype)
plan_loss_values, plan_err, plan_mask = laplacian_density_loss(target.float(), pred.float())
plan_loss = plan_loss_values.flatten(1).mean(dim=1)
flat_mask = plan_mask.flatten(1).float()
plan_mse = (plan_err.square().flatten(1) * flat_mask).sum(dim=1) / flat_mask.sum(dim=1).clamp_min(1.0)
weighted_plan_loss = (plan_loss_weight if "sample" in outputs else 1.0) * plan_loss
loss = weighted_plan_loss if loss is None else loss + weighted_plan_loss
terms["plan_loss"] = plan_loss.detach()
terms["plan_mse"] = plan_mse.detach()
weighted = weight * diffusion_loss
loss = weighted if loss is None else loss + weighted
terms[term] = diffusion_loss.detach()

for name, weight in (("plan", plan_loss_weight if "sample" in outputs else 1.0), ("action", action_loss_weight)):
if name not in outputs or (name == "plan" and name not in targets):
continue
pred = outputs[name]
target = targets[name].to(device=pred.device, dtype=pred.dtype if name == "plan" else torch.float32)
values, err, mask = laplacian_density_loss(target.float(), pred.float())
head_loss = values.flatten(1).mean(dim=1)
flat_mask = mask.flatten(1).float()
head_mse = (err.square().flatten(1) * flat_mask).sum(dim=1) / flat_mask.sum(dim=1).clamp_min(1.0)
weighted = weight * head_loss
loss = weighted if loss is None else loss + weighted
terms[f"{name}_loss"] = head_loss.detach()
terms[f"{name}_mse"] = head_mse.detach()

if loss is None:
raise RuntimeError("worldmodel produced no trainable outputs")
Expand All @@ -78,6 +93,7 @@ class WorldModelLoss(BaseLoss):
@dataclass(kw_only=True, slots=True)
class Config(BaseLoss.Config):
plan_loss_weight: float
action_loss_weight: float = 1.0

def __init__(self, config: Config, *, compile_config: CompileConfig | None = None):
plan_loss_weight = config.plan_loss_weight
Expand All @@ -90,6 +106,7 @@ def loss_fn(
outputs,
targets,
plan_loss_weight=plan_loss_weight,
action_loss_weight=config.action_loss_weight,
)

self.fn = loss_fn
Expand Down
Loading
Loading