from __future__ import annotations import torch from pino.pimt_model import DEFAULT_EMBEDDING_DIM, objective_targets_to_pyramid from pino.train import pad_trajectory_collate def _item(target_obj: torch.Tensor, *, n_mol: int = 2, t_steps: int = 49) -> dict: return { "tokens": torch.ones(n_mol, DEFAULT_EMBEDDING_DIM), "physics": torch.ones(t_steps, n_mol, 2), "target_obj": target_obj, "target_sub": torch.tensor([0.25, 0.25, 0.25, 0.25, 0.0, 0.5, 0.5]), } def test_objective_targets_to_pyramid_keeps_native_pyramid() -> None: targets = torch.rand(3, 138) normalized = objective_targets_to_pyramid(targets) assert normalized.shape == (3, 138) assert torch.equal(normalized, targets) def test_pad_trajectory_collate_normalizes_mixed_objective_target_shapes() -> None: pyramid = torch.zeros(3, 138) pyramid[0, 1] = 1.0 trajectory = torch.zeros(49, 138) trajectory[0, 2] = 0.4 trajectory[20, 3] = 0.7 trajectory[-1, 4] = 1.0 batch = pad_trajectory_collate([_item(pyramid), _item(trajectory, n_mol=3)]) assert batch["target_obj"].shape == (2, 3, 138) assert torch.equal(batch["target_obj"][0], pyramid) assert batch["target_obj"][1, 0, 2] == 0.4 assert batch["target_obj"][1, 1, 3] == 0.7 assert batch["target_obj"][1, 2, 4] == 1.0 def test_pad_trajectory_collate_normalizes_mixed_objective_target_shapes_when_legacy_first() -> None: trajectory = torch.zeros(49, 138) pyramid = torch.zeros(3, 138) batch = pad_trajectory_collate([_item(trajectory), _item(pyramid)]) assert batch["target_obj"].shape == (2, 3, 138)