|
| 1 | +# Copyright 2026 The LoongForge Authors. |
| 2 | +# SPDX-License-Identifier: Apache-2.0 |
| 3 | + |
| 4 | +import argparse |
| 5 | +import ast |
| 6 | +from pathlib import Path |
| 7 | +from types import SimpleNamespace |
| 8 | +from unittest.mock import Mock |
| 9 | + |
| 10 | +import pytest |
| 11 | + |
| 12 | + |
| 13 | +REPO_ROOT = Path(__file__).resolve().parents[1] |
| 14 | + |
| 15 | + |
| 16 | +def _load_functions(relative_path, names, namespace=None): |
| 17 | + """Load selected pure-Python functions without importing the GPU stack.""" |
| 18 | + path = REPO_ROOT / relative_path |
| 19 | + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) |
| 20 | + functions = [ |
| 21 | + node |
| 22 | + for node in tree.body |
| 23 | + if isinstance(node, ast.FunctionDef) and node.name in names |
| 24 | + ] |
| 25 | + assert {function.name for function in functions} == set(names) |
| 26 | + |
| 27 | + namespace = {} if namespace is None else namespace |
| 28 | + module = ast.Module(body=functions, type_ignores=[]) |
| 29 | + exec(compile(module, str(path), "exec"), namespace) |
| 30 | + return namespace |
| 31 | + |
| 32 | + |
| 33 | +@pytest.mark.parametrize( |
| 34 | + ("args", "expected"), |
| 35 | + [ |
| 36 | + ( |
| 37 | + SimpleNamespace(), |
| 38 | + {"shuffle_buffer_size": None, "max_samples_per_sequence": None}, |
| 39 | + ), |
| 40 | + ( |
| 41 | + SimpleNamespace( |
| 42 | + data_shuffle_buffer_size=1, |
| 43 | + data_max_samples_per_sequence=0, |
| 44 | + ), |
| 45 | + {"shuffle_buffer_size": None, "max_samples_per_sequence": None}, |
| 46 | + ), |
| 47 | + ( |
| 48 | + SimpleNamespace( |
| 49 | + data_shuffle_buffer_size=4096, |
| 50 | + data_max_samples_per_sequence=256, |
| 51 | + ), |
| 52 | + {"shuffle_buffer_size": 4096, "max_samples_per_sequence": 256}, |
| 53 | + ), |
| 54 | + ], |
| 55 | +) |
| 56 | +def test_energon_read_order_kwargs(args, expected): |
| 57 | + namespace = _load_functions( |
| 58 | + "loongforge/data/multimodal/dataloader_provider.py", |
| 59 | + {"_energon_read_order_kwargs"}, |
| 60 | + ) |
| 61 | + assert namespace["_energon_read_order_kwargs"](args) == expected |
| 62 | + |
| 63 | + |
| 64 | +@pytest.mark.parametrize( |
| 65 | + ("data_path", "expected_path"), |
| 66 | + [ |
| 67 | + (["dataset"], "dataset"), |
| 68 | + (["dataset-a", "dataset-b"], "/tmp/metadataset.yaml"), |
| 69 | + ], |
| 70 | +) |
| 71 | +def test_get_train_dataset_forwards_read_order_kwargs(data_path, expected_path): |
| 72 | + get_train_dataset = Mock(return_value="train-dataset") |
| 73 | + worker_config = object() |
| 74 | + energon = SimpleNamespace( |
| 75 | + WorkerConfig=Mock(return_value=worker_config), |
| 76 | + get_train_dataset=get_train_dataset, |
| 77 | + ) |
| 78 | + args = SimpleNamespace( |
| 79 | + data_path=data_path, |
| 80 | + micro_batch_size=2, |
| 81 | + num_workers=4, |
| 82 | + packing_buffer_size=10000, |
| 83 | + data_shuffle_buffer_size=4096, |
| 84 | + data_max_samples_per_sequence=256, |
| 85 | + rank=0, |
| 86 | + ) |
| 87 | + namespace = { |
| 88 | + "energon": energon, |
| 89 | + "parallel_state": SimpleNamespace( |
| 90 | + get_data_parallel_rank=lambda: 2, |
| 91 | + get_data_parallel_world_size=lambda: 8, |
| 92 | + get_data_parallel_group=lambda: "dp-group", |
| 93 | + ), |
| 94 | + "get_args": lambda: args, |
| 95 | + "get_blend_from_list": Mock( |
| 96 | + return_value=(["dataset-a", "dataset-b"], [0.25, 0.75]) |
| 97 | + ), |
| 98 | + "create_metadataset_yaml": Mock(return_value="/tmp/metadataset.yaml"), |
| 99 | + "print_error_handler": object(), |
| 100 | + "print_rank_0": Mock(), |
| 101 | + } |
| 102 | + namespace = _load_functions( |
| 103 | + "loongforge/data/multimodal/dataloader_provider.py", |
| 104 | + {"_energon_read_order_kwargs", "get_train_dataset"}, |
| 105 | + namespace, |
| 106 | + ) |
| 107 | + |
| 108 | + assert namespace["get_train_dataset"]("task-encoder") == "train-dataset" |
| 109 | + get_train_dataset.assert_called_once() |
| 110 | + path, = get_train_dataset.call_args.args |
| 111 | + kwargs = get_train_dataset.call_args.kwargs |
| 112 | + assert path == expected_path |
| 113 | + assert kwargs["shuffle_buffer_size"] == 4096 |
| 114 | + assert kwargs["max_samples_per_sequence"] == 256 |
| 115 | + assert kwargs["task_encoder"] == "task-encoder" |
| 116 | + assert kwargs["worker_config"] is worker_config |
| 117 | + |
| 118 | + |
| 119 | +def test_multimodal_argument_defaults_and_overrides(): |
| 120 | + class LanguageModelFamilies: |
| 121 | + @staticmethod |
| 122 | + def names(): |
| 123 | + return ["llama"] |
| 124 | + |
| 125 | + namespace = _load_functions( |
| 126 | + "loongforge/train/arguments.py", |
| 127 | + {"_add_extra_multimodal_args"}, |
| 128 | + { |
| 129 | + "get_support_model_archs": lambda values: values, |
| 130 | + "constants": SimpleNamespace(LanguageModelFamilies=LanguageModelFamilies), |
| 131 | + }, |
| 132 | + ) |
| 133 | + parser = namespace["_add_extra_multimodal_args"](argparse.ArgumentParser()) |
| 134 | + |
| 135 | + defaults = parser.parse_args([]) |
| 136 | + assert defaults.data_shuffle_buffer_size == 0 |
| 137 | + assert defaults.data_max_samples_per_sequence == 0 |
| 138 | + |
| 139 | + configured = parser.parse_args( |
| 140 | + [ |
| 141 | + "--data-shuffle-buffer-size", |
| 142 | + "4096", |
| 143 | + "--data-max-samples-per-sequence", |
| 144 | + "256", |
| 145 | + ] |
| 146 | + ) |
| 147 | + assert configured.data_shuffle_buffer_size == 4096 |
| 148 | + assert configured.data_max_samples_per_sequence == 256 |
0 commit comments