From 3da587a3dba56c33be1830120b5905b5b1cffb16 Mon Sep 17 00:00:00 2001 From: Walnutes Date: Wed, 19 Aug 2026 18:36:53 +0800 Subject: [PATCH] fix: unify reorder module, file and class name --- .../mixers/odm_dynamic_qwen_pt_full.yaml | 4 +- .../dynamic_saw_qwen_sft_full.yaml | 8 +-- .../train_full/reorder/saw_qwen_sft_full.yaml | 62 ++++++++++++++++ .../reorderers/saw_qwen_sft_full.yaml | 62 ---------------- .../doremi_step2_dynamic_qwen_pt_lora.yaml | 4 +- src/dataflex/cli.py | 72 +++++++++---------- src/dataflex/configs/components.yaml | 38 +++++----- src/dataflex/core/registry.py | 8 +-- src/dataflex/launcher.py | 70 +++++++++--------- src/dataflex/train/data/loader.py | 24 +++---- src/dataflex/train/reorder/__init__.py | 24 +++---- .../{base_reorderer.py => base_reorder.py} | 10 +-- ...ynamic_reorderer.py => dynamic_reorder.py} | 12 ++-- src/dataflex/train/reorder/score_provider.py | 4 +- ...{static_reorderer.py => static_reorder.py} | 14 ++-- src/dataflex/train/trainer/reorder_trainer.py | 42 +++++------ src/dataflex/utils/load_component.py | 37 +++++++++- src/dataflex/utils/selector_io.py | 6 +- 18 files changed, 266 insertions(+), 235 deletions(-) rename examples/train_full/{reorderers => reorder}/dynamic_saw_qwen_sft_full.yaml (59%) create mode 100644 examples/train_full/reorder/saw_qwen_sft_full.yaml delete mode 100644 examples/train_full/reorderers/saw_qwen_sft_full.yaml rename src/dataflex/train/reorder/{base_reorderer.py => base_reorder.py} (96%) rename src/dataflex/train/reorder/{dynamic_reorderer.py => dynamic_reorder.py} (95%) rename src/dataflex/train/reorder/{static_reorderer.py => static_reorder.py} (93%) diff --git a/examples/train_full/mixers/odm_dynamic_qwen_pt_full.yaml b/examples/train_full/mixers/odm_dynamic_qwen_pt_full.yaml index ecea294..5408d25 100644 --- a/examples/train_full/mixers/odm_dynamic_qwen_pt_full.yaml +++ b/examples/train_full/mixers/odm_dynamic_qwen_pt_full.yaml @@ -56,8 +56,8 @@ ddp_timeout: 180000000 ### dynamic_train - ODM: Online Data Mixing with Multi-Armed Bandits train_type: dynamic_mix components_cfg_file: src/dataflex/configs/components.yaml -component_name: odm # 使用ODM混合器 (Online Data Mixing with Exp3) -mixture_sample_rule: mixture # 初始采样规则,mixture为根据init_mixture_proportions比例混合 +component_name: odm # use ODM mixer (Online Data Mixing with Exp3) +mixture_sample_rule: mixture # initial sampling rule, mixture is according to the init_mixture_proportions ratio init_mixture_proportions: [0.5, 0.5] # initial weights warmup_step: 10 update_step: 10 diff --git a/examples/train_full/reorderers/dynamic_saw_qwen_sft_full.yaml b/examples/train_full/reorder/dynamic_saw_qwen_sft_full.yaml similarity index 59% rename from examples/train_full/reorderers/dynamic_saw_qwen_sft_full.yaml rename to examples/train_full/reorder/dynamic_saw_qwen_sft_full.yaml index 0692d3f..52e03ef 100644 --- a/examples/train_full/reorderers/dynamic_saw_qwen_sft_full.yaml +++ b/examples/train_full/reorder/dynamic_saw_qwen_sft_full.yaml @@ -9,7 +9,7 @@ finetuning_type: full deepspeed: examples/deepspeed/ds_z2_config.json ### dataset -# 动态排序不需要原始数据带分数字段:分数由当前模型在线算出来。 +# Dynamic reorder does not require original data to have a score field: the score is calculated online by the current model. dataset: alpaca_en_demo template: qwen cutoff_len: 2048 @@ -43,8 +43,8 @@ train_type: dynamic_reorder components_cfg_file: src/dataflex/configs/components.yaml component_name: dynamic_saw warmup_step: 10 -update_step: 50 # 每 50 步用当前模型重新给剩余样本打分并重排 +update_step: 50 # every 50 steps, use the current model to re-score and re-order the remaining samples update_times: 20 -# 开销提醒:每个 update 间隔要对样本池做一次前向。用 components.yaml 里的 -# score_params.max_samples 限制打分规模,或调大 reorder_every 降低打分频率。 +# Cost reminder: every update interval, the sample pool needs to do one forward pass. Use score_params.max_samples in components.yaml to limit the scoring scale, or increase reorder_every to reduce the scoring frequency. +# Use score_params.max_samples in components.yaml to limit the scoring scale, or increase reorder_every to reduce the scoring frequency. diff --git a/examples/train_full/reorder/saw_qwen_sft_full.yaml b/examples/train_full/reorder/saw_qwen_sft_full.yaml new file mode 100644 index 0000000..6e1d2f1 --- /dev/null +++ b/examples/train_full/reorder/saw_qwen_sft_full.yaml @@ -0,0 +1,62 @@ +### model +model_name_or_path: Qwen/Qwen3-1.7B-Base +trust_remote_code: true + +### method +stage: sft +do_train: true +finetuning_type: full +deepspeed: examples/deepspeed/ds_z2_config.json + +### dataset +# Original jsonl must contain a score field (see score_field in components.yaml). +# The paper uses "quality scores" like FineWeb-Edu's education score / QuRating score. +dataset: alpaca_en_demo +template: qwen +cutoff_len: 2048 +overwrite_cache: true +preprocessing_num_workers: 16 +dataloader_num_workers: 4 +seed: 42 + +# Constraints related to order: +# - val_size must be 0 (train_test_split will shuffle, shuffle the order); use eval_dataset for validation +# - mix_strategy must be concat (interleave_* will shuffle); use eval_dataset for validation +# - num_samples in dataset_info.json must not be set (it uses an unseeded random permutation) +val_size: 0 +mix_strategy: concat + +# packing will re-pack samples within a window of preprocessing_batch_size, which is equivalent to an implicit JIT of that width. +# To make a clean comparison, turn it off. +packing: false + +# It is recommended to use a separate tokenized_path for each reorder variant to avoid mixing up HF cache and avoid duplicate tokenization. +# tokenized_path: ../dataflex_saves/tokenized/alpaca_saw + +### output +output_dir: ../dataflex_saves/Qwen3-1.7B/reorder_saw +logging_steps: 10 +save_steps: 500 +plot_loss: true +overwrite_output_dir: true +report_to: none + +### train +per_device_train_batch_size: 4 +gradient_accumulation_steps: 8 +learning_rate: 2.0e-5 +num_train_epochs: 1.0 +lr_scheduler_type: cosine +warmup_ratio: 0.1 +bf16: true +ddp_timeout: 180000000 + +### Dataflex args +train_type: dynamic_reorder +components_cfg_file: src/dataflex/configs/components.yaml +component_name: saw # preset name in reorders in components.yaml + # optional: sorting(CL) / shuffle / segment / folding / zigzag / stair / saw +warmup_step: 10 +update_step: 50 +update_times: 20 # to run a full round, let warmup_step + update_step*update_times + # approximately equal to len(dataset) / (per_device_bs * grad_accum * world_size) diff --git a/examples/train_full/reorderers/saw_qwen_sft_full.yaml b/examples/train_full/reorderers/saw_qwen_sft_full.yaml deleted file mode 100644 index f7d6749..0000000 --- a/examples/train_full/reorderers/saw_qwen_sft_full.yaml +++ /dev/null @@ -1,62 +0,0 @@ -### model -model_name_or_path: Qwen/Qwen3-1.7B-Base -trust_remote_code: true - -### method -stage: sft -do_train: true -finetuning_type: full -deepspeed: examples/deepspeed/ds_z2_config.json - -### dataset -# 数据集的原始 jsonl 每行必须带一个分数字段(见 components.yaml 里的 score_field)。 -# 论文用的是 FineWeb-Edu 的教育分 / QuRating 分这类"质量分"。 -dataset: alpaca_en_demo -template: qwen -cutoff_len: 2048 -overwrite_cache: true -preprocessing_num_workers: 16 -dataloader_num_workers: 4 -seed: 42 - -# 顺序相关的约束: -# - val_size 必须为 0(train_test_split 会 shuffle,把顺序打乱);要做验证请用 eval_dataset -# - mix_strategy 保持 concat(interleave_* 会重排) -# - dataset_info.json 里对应条目的 num_samples 必须不设(它用的是未播种的随机置换) -val_size: 0 -mix_strategy: concat - -# packing 会在 preprocessing_batch_size 的窗口内按长度重组样本,相当于自带一个 -# 该宽度的隐式 JIT。想做干净对比就关掉它。 -packing: false - -# 每个排序变体建议用各自的 tokenized_path,避免 HF 缓存串味、也省去重复分词。 -# tokenized_path: ../dataflex_saves/tokenized/alpaca_saw - -### output -output_dir: ../dataflex_saves/Qwen3-1.7B/reorder_saw -logging_steps: 10 -save_steps: 500 -plot_loss: true -overwrite_output_dir: true -report_to: none - -### train -per_device_train_batch_size: 4 -gradient_accumulation_steps: 8 -learning_rate: 2.0e-5 -num_train_epochs: 1.0 -lr_scheduler_type: cosine -warmup_ratio: 0.1 -bf16: true -ddp_timeout: 180000000 - -### Dataflex args -train_type: dynamic_reorder -components_cfg_file: src/dataflex/configs/components.yaml -component_name: saw # components.yaml 的 reorderers 里的预设名 - # 可选:sorting(CL) / shuffle / segment / folding / zigzag / stair / saw -warmup_step: 10 -update_step: 50 -update_times: 20 # 想跑完整一轮就让 warmup_step + update_step*update_times - # 约等于 len(dataset) / (per_device_bs * grad_accum * world_size) diff --git a/examples/train_lora/mixers/doremi_step2_dynamic_qwen_pt_lora.yaml b/examples/train_lora/mixers/doremi_step2_dynamic_qwen_pt_lora.yaml index b8e72d1..0cf6ba3 100644 --- a/examples/train_lora/mixers/doremi_step2_dynamic_qwen_pt_lora.yaml +++ b/examples/train_lora/mixers/doremi_step2_dynamic_qwen_pt_lora.yaml @@ -57,8 +57,8 @@ ddp_timeout: 180000000 train_type: dynamic_mix components_cfg_file: src/dataflex/configs/components.yaml component_name: doremi -mixture_sample_rule: mixture # 初始采样规则,mixture为根据init_mixture_proportions比例混合(可动态调整),stratified为固定按源数据集大小比例分层,uniform为固定均匀分布 -init_mixture_proportions: [0.5, 0.5] # 对应初始的比例,可通过额外算法自行调整 +mixture_sample_rule: mixture # initial sampling rule, mixture is according to the init_mixture_proportions ratio (can be dynamically adjusted), stratified is fixed according to the source dataset size ratio, uniform is fixed uniform distribution +init_mixture_proportions: [0.5, 0.5] # corresponding initial proportions, can be adjusted by additional algorithms warmup_step: 100 update_step: 200 update_times: 3 diff --git a/src/dataflex/cli.py b/src/dataflex/cli.py index 691671c..a330f3b 100644 --- a/src/dataflex/cli.py +++ b/src/dataflex/cli.py @@ -75,23 +75,23 @@ def patch_trainer(train_type: str): TrainerCls = None if TrainerCls is not None: - # 1) 替换源头模块 + # 1) Replace source module tmod = importlib.import_module("llamafactory.train.sft.trainer") tmod.CustomSeq2SeqTrainer = TrainerCls - # 2) 替换包层 re-export + # 2) Replace package layer re-export sft_pkg = importlib.import_module("llamafactory.train.sft") setattr(sft_pkg, "CustomSeq2SeqTrainer", TrainerCls) - # 3) 替换 workflow 内部引用 + # 3) Replace workflow internal references wflow = importlib.import_module("llamafactory.train.sft.workflow") setattr(wflow, "CustomSeq2SeqTrainer", TrainerCls) - # 4) 替换 PT 训练器 + # 4) Replace PT trainer pt_tmod = importlib.import_module("llamafactory.train.pt.trainer") pt_tmod.CustomTrainer = TrainerCls - # 5) 替换 PT workflow 内部引用 + # 5) Replace PT workflow internal references pt_wflow = importlib.import_module("llamafactory.train.pt.workflow") setattr(pt_wflow, "CustomTrainer", TrainerCls) @@ -100,42 +100,42 @@ def patch_trainer(train_type: str): def patch_get_dataset(do_uncache_reload: bool = False): """ - 将 LlamaFactory 的 get_dataset 替换为 dataflex 版本。 - - 源头: llamafactory.data.loader.get_dataset -> dataflex.train.data.loader.get_dataset - - 包层 re-export: 覆盖 llamafactory.data.get_dataset(如有) - - 就地覆盖: 对已 from-import 的使用方(包含 workflow)直接改其全局符号 + Replace LlamaFactory's get_dataset with dataflex version. + - Source: llamafactory.data.loader.get_dataset -> dataflex.train.data.loader.get_dataset + - Package layer re-export: Overwrite llamafactory.data.get_dataset (if any) + - In-place overwrite: Directly modify the global symbol for already from-imported users (including workflow) Args: - do_uncache_reload: 为 True 时,会清理下游依赖缓存并预热导入,以确保后续 import 也拿到新函数。 - 默认为 False(与“就地打补丁”策略一致)。 + do_uncache_reload: When True, will clear downstream dependency cache and warm up imports to ensure subsequent imports also get the new function. + Default is False (consistent with "in-place patching" strategy). """ - # 1) 引入新实现 + # 1) Introduce new implementation from dataflex.train.data.loader import get_dataset as _new_get_dataset - # 2) 覆盖源头模块 + # 2) Overwrite source module data_loader_mod = importlib.import_module("llamafactory.data.loader") setattr(data_loader_mod, "get_dataset", _new_get_dataset) - # 3) 覆盖包层 re-export(若其它代码从包层 import) + # 3) Overwrite package layer re-export (if other code imports from package layer) data_pkg = importlib.import_module("llamafactory.data") setattr(data_pkg, "get_dataset", _new_get_dataset) - # 4) 就地覆盖已 from-import 的使用方(包含 workflow) + # 4) In-place overwrite already from-imported users (including workflow) wflow = importlib.import_module("llamafactory.train.sft.workflow") setattr(wflow, "get_dataset", _new_get_dataset) - # 5) 也要patch PT workflow + # 5) Also patch PT workflow pt_wflow = importlib.import_module("llamafactory.train.pt.workflow") setattr(pt_wflow, "get_dataset", _new_get_dataset) def patch_reorder_get_dataset(cfg): """ - 将 get_dataset 替换为"先按分数重排原始行、再做预处理"的版本。 + Replace get_dataset with the version that "first reorder raw rows by score, then preprocess". - 只有 apply_at == 'raw' 时才需要:分数字段在 align_dataset 里就被删掉了, - 而预处理不保 index(脏样本会被丢弃、packing 会合并行),所以直接重排原始 - 数据集,让顺序自然传递下去。apply_at == 'index' 时顺序在 trainer 里施加, - 数据加载流程无需改动。 + Only needed when apply_at == 'raw': the score field is removed in align_dataset, + and preprocessing does not preserve index (dirty samples are discarded, packing merges rows), + so we directly reorder the raw dataset to pass the order naturally. + When apply_at == 'index', the order is applied in trainer, and the data loading process remains unchanged. Returns: - bool: 是否真的打了补丁。 + bool: Whether the patch is actually applied. """ from dataflex.utils.load_component import load_component @@ -146,24 +146,24 @@ def patch_reorder_get_dataset(cfg): from dataflex.core.registry import REGISTRY from dataflex.train.data.loader import make_reorder_get_dataset - from dataflex.train.reorder import resolve_reorderer_kind # also registers the reorderers + from dataflex.train.reorder import resolve_reorder_kind # also registers the reorders - params = load_component('reorderers', cfg_file, name, runtime_vars={}) - kind = resolve_reorderer_kind(name, params) + params = load_component('reorders', cfg_file, name, runtime_vars={}) + kind = resolve_reorder_kind(name, params) - # 只有"静态 + 在原始行上重排"才需要改数据加载。动态排序的分数来自当前模型, - # 顺序必然是在 trainer 里按 dataset index 施加的。 + # Only "static + reorder on raw rows" needs to modify data loading. The dynamic sorting scores come from the current model, + # the order must be applied in trainer by dataset index. if kind != 'static' or params.get('apply_at', 'raw') != 'raw': - print(f"[PatchReorder] reorderer '{name}' orders by dataset index; dataset loading left untouched.") + print(f"[PatchReorder] reorder '{name}' orders by dataset index; dataset loading left untouched.") return False - def reorderer_factory(): - return REGISTRY.build('reorderer', kind, runtime={}, cfg=params) + def reorder_factory(): + return REGISTRY.build('reorder', kind, runtime={}, cfg=params) - _new_get_dataset = make_reorder_get_dataset(reorderer_factory) + _new_get_dataset = make_reorder_get_dataset(reorder_factory) - # 与 patch_get_dataset 同样的四处覆盖:源头模块、包层 re-export、以及 - # sft/pt 两个 workflow 里已经 from-import 过的全局符号。 + # Same four patches as patch_get_dataset: source module, package layer re-export, and + # already from-imported global symbols in sft/pt workflows. data_loader_mod = importlib.import_module("llamafactory.data.loader") setattr(data_loader_mod, "get_dataset", _new_get_dataset) data_pkg = importlib.import_module("llamafactory.data") @@ -173,7 +173,7 @@ def reorderer_factory(): pt_wflow = importlib.import_module("llamafactory.train.pt.workflow") setattr(pt_wflow, "get_dataset", _new_get_dataset) - print(f"[PatchReorder] reorderer '{name}' will permute raw rows before preprocessing.") + print(f"[PatchReorder] reorder '{name}' will permute raw rows before preprocessing.") return True def read_args(): @@ -184,7 +184,7 @@ def read_args(): dict_config = OmegaConf.load(Path(file_path).absolute()) cfg = OmegaConf.merge(dict_config, override_config) else: - cfg = OmegaConf.create({}) # CLI 直接传参时 + cfg = OmegaConf.create({}) # When passing CLI arguments directly return OmegaConf.to_container(cfg) @@ -211,7 +211,7 @@ def print_welcome(): def main(): command = sys.argv.pop(1) if command == "version": - # 只打印版本和欢迎 + # Only print version and welcome print_welcome() return elif command != 'train': diff --git a/src/dataflex/configs/components.yaml b/src/dataflex/configs/components.yaml index 1193c19..f91bf5d 100644 --- a/src/dataflex/configs/components.yaml +++ b/src/dataflex/configs/components.yaml @@ -170,26 +170,26 @@ weighters: adapt: name: adapt params: - tau: 1.0 # 温度,越小权重区分越锐利 - refresh_interval: 50 # 每多少步用当前模型刷新一次 anchor 向量 - anchor_batch_size: 8 # 计算句向量时的前向 batch 大小 - clip: null # 可选权重上限,防梯度爆炸 + tau: 1.0 # Temperature, smaller values make weight distinction sharper + refresh_interval: 50 # Every how many steps to refresh the anchor vector with the current model + anchor_batch_size: 8 # Forward batch size for computing sentence embeddings + clip: null # Optional weight upper bound, to prevent gradient explosion joint_update_aware: name: joint_update_aware params: - # 求解 max_w s^T w - (beta / 2) w^T S w + tau * H(w) - # S 为归一化 embedding 的余弦 Gram 矩阵,s_i = 为对 anchor 均值 u 的对齐度 - beta: 0.1 # 交互/冗余强度 - tau: 0.05 # 熵温度(需 > 0),越小权重越锐利 - fixed_point_iters: 5 # 阻尼定点迭代次数 - damping: 1.0 # 阻尼系数 rho ∈ (0, 1],1.0 表示不阻尼 - target_update_step: 50 # 每多少步用 eval 集刷新一次目标向量 - target_batch_size: 1 # 计算目标向量时的前向 batch 大小 - target_num_batches: 1 # 每次刷新平均多少个 eval batch;设为 0 表示遍历整个 eval 集 - embed_normalize: true # 是否对 embedding 做 L2 归一化 - pooling: last_token # 句向量池化方式:last_token / mean_pool - embed_layer: -1 # 取哪一层 hidden state,-1 为最后一层 + # Solve max_w s^T w - (beta / 2) w^T S w + tau * H(w) + # S is the cosine Gram matrix of normalized embeddings, s_i = is the alignment degree of the anchor mean u + beta: 0.1 # Interaction/redundancy strength + tau: 0.05 # Entropy temperature (must be > 0), smaller values make weights sharper + fixed_point_iters: 5 # Damping fixed point iterations + damping: 1.0 # Damping coefficient rho ∈ (0, 1], 1.0 means no damping + target_update_step: 50 # Every how many steps to refresh the target vector with the eval set + target_batch_size: 1 # Forward batch size for computing target embeddings + target_num_batches: 1 # Average how many eval batches to refresh per target; set to 0 to iterate through the entire eval set + embed_normalize: true # Whether to normalize embeddings with L2 norm + pooling: last_token # Sentence vector pooling: last_token / mean_pool + embed_layer: -1 # Which hidden state layer to take, -1 means the last layer objective_mode: full # full / align_only / diverse_only / uniform custom: @@ -197,10 +197,10 @@ weighters: params: strategy: uniform -reorderers: - # Data reorderers change sample order only; the number of samples is untouched. +reorders: + # Data reorders change sample order only; the number of samples is untouched. # - # Every reorderer is defined along three axes: + # Every reorder is defined along three axes: # pattern : one of seven orderings, covering the four guidances of the paper # [shuffle, sorting, folding, zigzag, segment, stair, saw] # score_source : where the scores come from [precomputed, cached_selection, model_loss] diff --git a/src/dataflex/core/registry.py b/src/dataflex/core/registry.py index 7787958..62b55a2 100644 --- a/src/dataflex/core/registry.py +++ b/src/dataflex/core/registry.py @@ -20,14 +20,14 @@ def get(self, kind: str, name: str) -> Type: def build(self, kind: str, name: str, *, runtime: Dict[str, Any], cfg: Optional[Dict[str, Any]] = None): cls = self.get(kind, name) cfg = cfg or {} - merged = {**cfg, **runtime} # 运行期依赖优先 + merged = {**cfg, **runtime} # Runtime dependencies take precedence sig = inspect.signature(cls.__init__) - accepted = {p.name for p in list(sig.parameters.values())[1:]} # 跳过 self - filtered = {k: v for k, v in merged.items() if k in accepted} # 只喂需要的 + accepted = {p.name for p in list(sig.parameters.values())[1:]} # Skip self + filtered = {k: v for k, v in merged.items() if k in accepted} # Only feed needed return cls(**filtered) REGISTRY = Registry() def register_selector(name: str): return REGISTRY.register("selector", name) def register_mixer(name: str): return REGISTRY.register("mixer", name) def register_weighter(name: str): return REGISTRY.register("weighter", name) -def register_reorderer(name: str): return REGISTRY.register("reorderer", name) +def register_reorder(name: str): return REGISTRY.register("reorder", name) diff --git a/src/dataflex/launcher.py b/src/dataflex/launcher.py index 78d5d7d..f367b60 100644 --- a/src/dataflex/launcher.py +++ b/src/dataflex/launcher.py @@ -72,23 +72,23 @@ def patch_trainer(train_type: str): TrainerCls = None if TrainerCls is not None: - # 1) 替换源头模块 + # 1) Replace source module tmod = importlib.import_module("llamafactory.train.sft.trainer") tmod.CustomSeq2SeqTrainer = TrainerCls - # 2) 替换包层 re-export + # 2) Replace package layer re-export sft_pkg = importlib.import_module("llamafactory.train.sft") setattr(sft_pkg, "CustomSeq2SeqTrainer", TrainerCls) - # 3) 替换 workflow 内部引用 + # 3) Replace workflow internal references wflow = importlib.import_module("llamafactory.train.sft.workflow") setattr(wflow, "CustomSeq2SeqTrainer", TrainerCls) - # 4) 替换 PT 训练器 + # 4) Replace PT trainer pt_tmod = importlib.import_module("llamafactory.train.pt.trainer") pt_tmod.CustomTrainer = TrainerCls - # 5) 替换 PT workflow 内部引用 + # 5) Replace PT workflow internal references pt_wflow = importlib.import_module("llamafactory.train.pt.workflow") setattr(pt_wflow, "CustomTrainer", TrainerCls) @@ -96,42 +96,42 @@ def patch_trainer(train_type: str): def patch_get_dataset(do_uncache_reload: bool = False): """ - 将 LlamaFactory 的 get_dataset 替换为 dataflex 版本。 - - 源头: llamafactory.data.loader.get_dataset -> dataflex.train.data.loader.get_dataset - - 包层 re-export: 覆盖 llamafactory.data.get_dataset(如有) - - 就地覆盖: 对已 from-import 的使用方(包含 workflow)直接改其全局符号 + Replace LlamaFactory's get_dataset with dataflex version. + - Source: llamafactory.data.loader.get_dataset -> dataflex.train.data.loader.get_dataset + - Package layer re-export: Overwrite llamafactory.data.get_dataset (if any) + - In-place overwrite: Directly modify the global symbol for already from-imported users (including workflow) Args: - do_uncache_reload: 为 True 时,会清理下游依赖缓存并预热导入,以确保后续 import 也拿到新函数。 - 默认为 False(与“就地打补丁”策略一致)。 + do_uncache_reload: When True, will clear downstream dependency cache and warm up imports to ensure subsequent imports also get the new function. + Default is False (consistent with "in-place patching" strategy). """ - # 1) 引入新实现 + # 1) Introduce new implementation from dataflex.train.data.loader import get_dataset as _new_get_dataset - # 2) 覆盖源头模块 + # 2) Overwrite source module data_loader_mod = importlib.import_module("llamafactory.data.loader") setattr(data_loader_mod, "get_dataset", _new_get_dataset) - # 3) 覆盖包层 re-export(若其它代码从包层 import) + # 3) Overwrite package layer re-export (if other code imports from package layer) data_pkg = importlib.import_module("llamafactory.data") setattr(data_pkg, "get_dataset", _new_get_dataset) - # 4) 就地覆盖已 from-import 的使用方(包含 workflow) + # 4) In-place overwrite already from-imported users (including workflow) wflow = importlib.import_module("llamafactory.train.sft.workflow") setattr(wflow, "get_dataset", _new_get_dataset) - # 5) 也要patch PT workflow + # 5) Also patch PT workflow pt_wflow = importlib.import_module("llamafactory.train.pt.workflow") setattr(pt_wflow, "get_dataset", _new_get_dataset) def patch_reorder_get_dataset(cfg): """ - 将 get_dataset 替换为"先按分数重排原始行、再做预处理"的版本。 + Replace get_dataset with the version that "first reorder raw rows by score, then preprocess". - 只有 apply_at == 'raw' 时才需要:分数字段在 align_dataset 里就被删掉了, - 而预处理不保 index(脏样本会被丢弃、packing 会合并行),所以直接重排原始 - 数据集,让顺序自然传递下去。apply_at == 'index' 时顺序在 trainer 里施加, - 数据加载流程无需改动。 + Only needed when apply_at == 'raw': the score field is removed in align_dataset, + and preprocessing does not preserve index (dirty samples are discarded, packing merges rows), + so we directly reorder the raw dataset to pass the order naturally. + When apply_at == 'index', the order is applied in trainer, and the data loading process remains unchanged. Returns: - bool: 是否真的打了补丁。 + bool: Whether the patch is actually applied. """ from dataflex.utils.load_component import load_component @@ -142,24 +142,24 @@ def patch_reorder_get_dataset(cfg): from dataflex.core.registry import REGISTRY from dataflex.train.data.loader import make_reorder_get_dataset - from dataflex.train.reorder import resolve_reorderer_kind # also registers the reorderers + from dataflex.train.reorder import resolve_reorder_kind # also registers the reorders - params = load_component('reorderers', cfg_file, name, runtime_vars={}) - kind = resolve_reorderer_kind(name, params) + params = load_component('reorders', cfg_file, name, runtime_vars={}) + kind = resolve_reorder_kind(name, params) - # 只有"静态 + 在原始行上重排"才需要改数据加载。动态排序的分数来自当前模型, - # 顺序必然是在 trainer 里按 dataset index 施加的。 + # Only "static + reorder on raw rows" needs to modify data loading. The dynamic sorting scores come from the current model, + # the order must be applied in trainer by dataset index. if kind != 'static' or params.get('apply_at', 'raw') != 'raw': - print(f"[PatchReorder] reorderer '{name}' orders by dataset index; dataset loading left untouched.") + print(f"[PatchReorder] reorder '{name}' orders by dataset index; dataset loading left untouched.") return False - def reorderer_factory(): - return REGISTRY.build('reorderer', kind, runtime={}, cfg=params) + def reorder_factory(): + return REGISTRY.build('reorder', kind, runtime={}, cfg=params) - _new_get_dataset = make_reorder_get_dataset(reorderer_factory) + _new_get_dataset = make_reorder_get_dataset(reorder_factory) - # 与 patch_get_dataset 同样的四处覆盖:源头模块、包层 re-export、以及 - # sft/pt 两个 workflow 里已经 from-import 过的全局符号。 + # Same four patches as patch_get_dataset: source module, package layer re-export, and + # already from-imported global symbols in sft/pt workflows. data_loader_mod = importlib.import_module("llamafactory.data.loader") setattr(data_loader_mod, "get_dataset", _new_get_dataset) data_pkg = importlib.import_module("llamafactory.data") @@ -169,7 +169,7 @@ def reorderer_factory(): pt_wflow = importlib.import_module("llamafactory.train.pt.workflow") setattr(pt_wflow, "get_dataset", _new_get_dataset) - print(f"[PatchReorder] reorderer '{name}' will permute raw rows before preprocessing.") + print(f"[PatchReorder] reorder '{name}' will permute raw rows before preprocessing.") return True @@ -181,7 +181,7 @@ def read_args(): dict_config = OmegaConf.load(Path(file_path).absolute()) cfg = OmegaConf.merge(dict_config, override_config) else: - cfg = OmegaConf.create({}) # CLI 直接传参时 + cfg = OmegaConf.create({}) # When passing CLI arguments directly return OmegaConf.to_container(cfg) diff --git a/src/dataflex/train/data/loader.py b/src/dataflex/train/data/loader.py index f96fc60..b75f82f 100644 --- a/src/dataflex/train/data/loader.py +++ b/src/dataflex/train/data/loader.py @@ -117,7 +117,7 @@ def get_dataset( sizes_str = {name: len(ds) for name, ds in per_source_pp.items()} logger.info_rank0(f"[Dataflex] Per-source preprocessed sizes: {sizes_str}") - # 打印初始比例配置 + # Print initial proportion configuration logger.info_rank0(f"[Dataflex] sample_rule={data_args.mixture_sample_rule} | " f"proportions={data_args.init_mixture_proportions} | " f"seed={training_args.seed}") @@ -152,9 +152,9 @@ def get_dataset( logger.info_rank0(f"[Dataflex] Mixer eval: '{eval_name}' -> domain '{domain}' ({len(eval_ds)} samples)") manager.mixer_eval_datasets = mixer_eval_by_domain - # 可选:把 manager 留给外部(方便在 callback 里重建) - # 例如附在 dataset_module 上(Trainer 不会用到这个字段) - dataset_module["train_dataset"] = None # 先占位,trainer里会rebuild + # Optional: expose manager to external code (for callback reconstruction) + # e.g. attached to dataset_module (Trainer doesn't use this field) + dataset_module["train_dataset"] = None # Placeholder, trainer will rebuild dataset_module["mixture_manager"] = manager logger.info_rank0("[Dataflex] Exposed mixture_manager for runtime re-mixing.") @@ -211,13 +211,13 @@ def _concat_raw_scores(captured: List[Optional[np.ndarray]], score_field: str) - if not captured or any(part is None for part in captured): raise ValueError( f"[Dataflex][Reorder] score field '{score_field}' is missing from at least one dataset. " - f"Either add it to every source, or switch the reorderer to " + f"Either add it to every source, or switch the reorder to " f"`apply_at: index` with an explicit `score_path`." ) return np.concatenate(captured, axis=0) -def make_reorder_get_dataset(reorderer_factory): +def make_reorder_get_dataset(reorder_factory): """Build a `get_dataset` that permutes raw rows before tokenization. Why here and not in the trainer: the score lives in the raw JSONL and is @@ -228,7 +228,7 @@ def make_reorder_get_dataset(reorderer_factory): relative position. Args: - reorderer_factory: zero-arg callable returning a reorderer exposing + reorder_factory: zero-arg callable returning a reorder exposing `order_rows(scores) -> permutation` and a `score_params` dict. """ @@ -255,8 +255,8 @@ def reorder_get_dataset( if data_args.streaming: raise ValueError("[Dataflex][Reorder] reordering requires `streaming: false`.") - reorderer = reorderer_factory() - score_field = reorderer.score_params.get("score_field", "score") + reorder = reorder_factory() + score_field = reorder.score_params.get("score_field", "score") with training_args.main_process_first(desc="load dataset", local=(not data_args.data_shared_file_system)): with _capture_raw_scores(score_field) as captured: @@ -272,7 +272,7 @@ def reorder_get_dataset( ) if dataset is not None: - score_path = reorderer.score_params.get("score_path") + score_path = reorder.score_params.get("score_path") if score_path: from ..reorder.score_provider import PrecomputedScoreProvider @@ -292,10 +292,10 @@ def reorder_get_dataset( f"carries '{score_field}'." ) - permutation = reorderer.order_rows(scores) + permutation = reorder.order_rows(scores) dataset = dataset.select(permutation) logger.info_rank0( - f"[Dataflex][Reorder] applied '{reorderer.pattern}' to {len(permutation)} raw rows " + f"[Dataflex][Reorder] applied '{reorder.pattern}' to {len(permutation)} raw rows " f"before preprocessing." ) diff --git a/src/dataflex/train/reorder/__init__.py b/src/dataflex/train/reorder/__init__.py index d0d0a74..95a4a43 100644 --- a/src/dataflex/train/reorder/__init__.py +++ b/src/dataflex/train/reorder/__init__.py @@ -1,12 +1,12 @@ """Data ordering components. -Importing this package is what registers the reorderers, so every reorderer must +Importing this package is what registers the reorders, so every reorder must be imported here. Registry names must be unique or `REGISTRY.register` raises at import time. """ -from .base_reorderer import POLARITIES, Reorderer -from .dynamic_reorderer import DynamicReorderer +from .base_reorder import POLARITIES, Reorder +from .dynamic_reorder import DynamicReorder from .patterns import PATTERNS, apply_pattern from .score_provider import ( CachedSelectionScoreProvider, @@ -15,17 +15,17 @@ ScoreProvider, build_score_provider, ) -from .static_reorderer import StaticReorderer +from .static_reorder import StaticReorder -#: Registered reorderer classes. Unlike the other families, a reorderer preset in +#: Registered reorder classes. Unlike the other families, a reorder preset in #: components.yaml is named after the *ordering* it produces (`saw`, `folding`, #: ...) rather than after its class, because the ordering is chosen by config. #: `kind` is what maps a preset onto one of these. KINDS = ("static", "dynamic") -def resolve_reorderer_kind(component_name: str, params: dict) -> str: - """Decide which reorderer class a components.yaml preset refers to. +def resolve_reorder_kind(component_name: str, params: dict) -> str: + """Decide which reorder class a components.yaml preset refers to. Falls back to the preset name so a preset named directly after a class still works without a `kind`, matching how selectors/mixers/weighters behave. @@ -33,7 +33,7 @@ def resolve_reorderer_kind(component_name: str, params: dict) -> str: kind = params.pop("kind", None) or component_name if kind not in KINDS: raise ValueError( - f"reorderer preset '{component_name}' resolves to kind '{kind}', which is not registered. " + f"reorder preset '{component_name}' resolves to kind '{kind}', which is not registered. " f"Set `kind` to one of {list(KINDS)} in its params block." ) return kind @@ -43,14 +43,14 @@ def resolve_reorderer_kind(component_name: str, params: dict) -> str: "POLARITIES", "PATTERNS", "KINDS", - "Reorderer", - "StaticReorderer", - "DynamicReorderer", + "Reorder", + "StaticReorder", + "DynamicReorder", "ScoreProvider", "PrecomputedScoreProvider", "CachedSelectionScoreProvider", "ModelLossScoreProvider", "build_score_provider", "apply_pattern", - "resolve_reorderer_kind", + "resolve_reorder_kind", ] diff --git a/src/dataflex/train/reorder/base_reorderer.py b/src/dataflex/train/reorder/base_reorder.py similarity index 96% rename from src/dataflex/train/reorder/base_reorderer.py rename to src/dataflex/train/reorder/base_reorder.py index 7675ea3..5efb0b5 100644 --- a/src/dataflex/train/reorder/base_reorderer.py +++ b/src/dataflex/train/reorder/base_reorder.py @@ -29,10 +29,10 @@ } -class Reorderer(ABC): +class Reorder(ABC): """Base class for data ordering components. - A reorderer turns a set of *positions* plus their scores into a permutation. + A reorder turns a set of *positions* plus their scores into a permutation. Positions are opaque integers whose meaning depends on the path: raw dataset rows for the static path, indices into the preprocessed ``train_dataset`` for the dynamic path. @@ -84,7 +84,7 @@ def __init__( # # These exist so reorder can later be chained with select/mix/weight into a # single data pipeline without reworking the component. They are all no-ops - # by default, so a standalone reorderer behaves exactly as if they were + # by default, so a standalone reorder behaves exactly as if they were # absent. # ------------------------------------------------------------------ @@ -92,8 +92,8 @@ def set_candidate_pool(self, indices: Optional[Sequence[int]]) -> None: """Restrict ordering to a subset of positions. This is the seam for chaining after a selector: the selector decides - *which* samples survive, the reorderer then decides their order. When - unset, the reorderer orders everything it is given. + *which* samples survive, this component then decides their order. When + unset, it orders everything it is given. """ self._candidate_pool = list(indices) if indices is not None else None diff --git a/src/dataflex/train/reorder/dynamic_reorderer.py b/src/dataflex/train/reorder/dynamic_reorder.py similarity index 95% rename from src/dataflex/train/reorder/dynamic_reorderer.py rename to src/dataflex/train/reorder/dynamic_reorder.py index ad13e87..bcbdef2 100644 --- a/src/dataflex/train/reorder/dynamic_reorderer.py +++ b/src/dataflex/train/reorder/dynamic_reorder.py @@ -3,15 +3,15 @@ import numpy as np import torch.distributed as dist -from dataflex.core.registry import register_reorderer +from dataflex.core.registry import register_reorder from dataflex.utils.logging import logger -from .base_reorderer import Reorderer +from .base_reorder import Reorder from .score_provider import build_score_provider -@register_reorderer("dynamic") -class DynamicReorderer(Reorderer): +@register_reorder("dynamic") +class DynamicReorder(Reorder): """Re-score and re-order the remaining pool as training progresses. The static path fixes the whole curriculum before the first step, so the @@ -71,7 +71,7 @@ def __init__( self._call_count = 0 logger.info( - f"[Dataflex][Reorder] DynamicReorderer({self.describe()}, score_source={score_source}, " + f"[Dataflex][Reorder] DynamicReorder({self.describe()}, score_source={score_source}, " f"reorder_every={self.reorder_every}, consume_once={self.consume_once})" ) @@ -93,7 +93,7 @@ def _ensure_pool(self) -> List[int]: pool = self.get_candidate_pool() if pool is None: if self.dataset is None: - raise ValueError("DynamicReorderer needs a dataset or an explicit candidate pool") + raise ValueError("DynamicReorder needs a dataset or an explicit candidate pool") pool = list(range(len(self.dataset))) self._pool = list(pool) return self._pool diff --git a/src/dataflex/train/reorder/score_provider.py b/src/dataflex/train/reorder/score_provider.py index d4c3e36..15b7dfe 100644 --- a/src/dataflex/train/reorder/score_provider.py +++ b/src/dataflex/train/reorder/score_provider.py @@ -15,7 +15,7 @@ argument expressed with DataFlex's own artifacts. Note on polarity: these return raw scores. Interpreting whether high means good -is `Reorderer.polarity`, not the provider's job. +is `Reorder.polarity`, not the provider's job. """ import json @@ -158,7 +158,7 @@ class CachedSelectionScoreProvider(PrecomputedScoreProvider): """Reuse a score a selector already computed and cached. `save_selection` writes `{"indices": [...], "metric": {"loss": [...]}}`, so a - reorderer can consume a selector's work instead of paying for a second + reorder can consume a selector's work instead of paying for a second scoring pass. Positions absent from the cache get `fill_value`. """ diff --git a/src/dataflex/train/reorder/static_reorderer.py b/src/dataflex/train/reorder/static_reorder.py similarity index 93% rename from src/dataflex/train/reorder/static_reorderer.py rename to src/dataflex/train/reorder/static_reorder.py index 6db6b6a..0dc2802 100644 --- a/src/dataflex/train/reorder/static_reorderer.py +++ b/src/dataflex/train/reorder/static_reorder.py @@ -1,14 +1,14 @@ from typing import List, Optional, Sequence -from dataflex.core.registry import register_reorderer +from dataflex.core.registry import register_reorder from dataflex.utils.logging import logger -from .base_reorderer import Reorderer +from .base_reorder import Reorder from .score_provider import build_score_provider -@register_reorderer("static") -class StaticReorderer(Reorderer): +@register_reorder("static") +class StaticReorder(Reorder): """Order the dataset once, from scores that do not depend on the model. This is the faithful reproduction of the paper: a single global permutation @@ -68,7 +68,7 @@ def __init__( self._cursor = 0 logger.info( - f"[Dataflex][Reorder] StaticReorderer({self.describe()}, " + f"[Dataflex][Reorder] StaticReorder({self.describe()}, " f"score_source={score_source}, apply_at={apply_at})" ) @@ -95,7 +95,7 @@ def _ensure_order(self, model=None, step_id: int = 0) -> List[int]: pool = self.get_candidate_pool() if pool is None: if self.dataset is None: - raise ValueError("StaticReorderer needs a dataset or an explicit candidate pool") + raise ValueError("StaticReorder needs a dataset or an explicit candidate pool") pool = list(range(len(self.dataset))) if self.apply_at == "raw": @@ -119,7 +119,7 @@ def _ensure_order(self, model=None, step_id: int = 0) -> List[int]: if provider.is_dynamic: logger.warning( f"[Dataflex][Reorder] score source '{self.score_source}' depends on the model but " - f"StaticReorderer scores only once, at the first update. Use the 'dynamic' reorderer " + f"StaticReorder scores only once, at the first update. Use the 'dynamic' reorder " f"to re-score as training progresses." ) diff --git a/src/dataflex/train/trainer/reorder_trainer.py b/src/dataflex/train/trainer/reorder_trainer.py index 1a25f0b..0d021ec 100644 --- a/src/dataflex/train/trainer/reorder_trainer.py +++ b/src/dataflex/train/trainer/reorder_trainer.py @@ -7,15 +7,15 @@ from dataflex.utils.load_component import load_component from dataflex.utils.logging import logger -from dataflex.train.reorder import resolve_reorderer_kind # also registers the reorderers +from dataflex.train.reorder import resolve_reorder_kind # also registers the reorders from .select_trainer import SelectTrainer -class _ReordererAsSelector: - """Adapt a reorderer to the index-provider protocol `SelectTrainer` expects. +class _ReorderAsSelector: + """Adapt a reorder to the index-provider protocol `SelectTrainer` expects. - `SelectTrainer` already does exactly what a reorderer needs: every + `SelectTrainer` already does exactly what a reorder needs: every `update_step` steps it asks a component for a list of indices, wraps them in `torch.utils.data.Subset` (which respects list order) and swaps the iterator. Reusing that loop rather than copying a fourth ~550-line @@ -26,27 +26,27 @@ class _ReordererAsSelector: and `select` returns an ordered chunk rather than a scored subset. """ - def __init__(self, reorderer, accelerator=None): - self.reorderer = reorderer + def __init__(self, reorder, accelerator=None): + self.reorder = reorder self.accelerator = accelerator self.data_collator = None # assigned by get_train_dataloader def warmup(self, num_samples: int, replacement: bool = False) -> List[int]: # The base Selector.warmup samples randomly with replacement, which for a # score-ordered run would start the curriculum in the middle. The - # reorderer decides instead: the head of the ordering when scores are + # reorder decides instead: the head of the ordering when scores are # model-independent, a random draw when they are not. - return self.reorderer.warmup_indices(num_samples) + return self.reorder.warmup_indices(num_samples) def select(self, model, step_id: int, num_samples: int, **kwargs) -> List[int]: - return self.reorderer.next_indices(model=model, step_id=step_id, num_samples=num_samples, **kwargs) + return self.reorder.next_indices(model=model, step_id=step_id, num_samples=num_samples, **kwargs) def __getattr__(self, item): - # Forward anything else (e.g. observe) to the reorderer. Guarded so an + # Forward anything else (e.g. observe) to the reorder. Guarded so an # access before __init__ finishes raises instead of recursing forever. - if item == "reorderer": + if item == "reorder": raise AttributeError(item) - return getattr(self.reorderer, item) + return getattr(self.reorder, item) class ReorderTrainer(SelectTrainer): @@ -69,8 +69,8 @@ def __init__(self, finetuning_args, processor=None, gen_kwargs=None, model_args= ) name = finetuning_args.component_name - params = load_component("reorderers", finetuning_args.components_cfg_file, name, runtime_vars={}) - kind = resolve_reorderer_kind(name, params) + params = load_component("reorders", finetuning_args.components_cfg_file, name, runtime_vars={}) + kind = resolve_reorder_kind(name, params) runtime = dict( dataset=self.train_dataset, @@ -78,20 +78,20 @@ def __init__(self, finetuning_args, processor=None, gen_kwargs=None, model_args= accelerator=self.accelerator, data_collator=self.data_collator, ) - self.reorderer = REGISTRY.build("reorderer", kind, runtime=runtime, cfg=params) + self.reorder = REGISTRY.build("reorder", kind, runtime=runtime, cfg=params) # SelectTrainer's loop calls `self.selector`; the adapter makes the - # reorderer answer to that protocol without editing select_trainer.py. - self.selector = _ReordererAsSelector(self.reorderer, accelerator=self.accelerator) + # reorder answer to that protocol without editing select_trainer.py. + self.selector = _ReorderAsSelector(self.reorder, accelerator=self.accelerator) - logger.info(f"[ReorderTrainer] reorderer={name} (kind={kind}), params={params}") + logger.info(f"[ReorderTrainer] reorder={name} (kind={kind}), params={params}") logger.info("[Dataflex] ReorderTrainer initialized") @override def _get_train_sampler(self, train_dataset=None) -> Optional[torch.utils.data.Sampler]: """Always sequential. - The reorderer's output order *is* the curriculum, so any shuffling + The reorder's output order *is* the curriculum, so any shuffling sampler silently discards it and the run degrades into the random baseline while still looking correct. This is forced rather than left to `disable_shuffling` because that failure is invisible in the logs. @@ -113,11 +113,11 @@ def _get_train_sampler(self, train_dataset=None) -> Optional[torch.utils.data.Sa @override def _maybe_log_save_evaluate(self, tr_loss, grad_norm, model, trial, epoch, ignore_keys_for_eval, *args, **kwargs): - # Feed the training signal back to the reorderer. Nothing consumes it + # Feed the training signal back to the reorder. Nothing consumes it # yet; this is the seam for an adaptive controller that would react to # gradient-norm spikes at cycle boundaries. try: - self.reorderer.observe( + self.reorder.observe( global_step=self.state.global_step, grad_norm=float(grad_norm) if grad_norm is not None else None, learning_rate=self._get_learning_rate(), diff --git a/src/dataflex/utils/load_component.py b/src/dataflex/utils/load_component.py index 7cc25d8..c9fe6b2 100644 --- a/src/dataflex/utils/load_component.py +++ b/src/dataflex/utils/load_component.py @@ -1,16 +1,47 @@ import yaml from typing import Dict, Any, Optional +from dataflex.utils.logging import logger + +#: Bucket names that were renamed, mapped to the older spellings still accepted. +#: The reorder family used to be called "reorderer", so a config written before +#: the rename says `reorderers:`. Reading it keeps working; only a warning marks +#: it as deprecated. +_BUCKET_ALIASES: Dict[str, tuple] = { + "reorders": ("reorderers",), +} + +_warned_aliases = set() + + +def _resolve_bucket(root: Dict[str, Any], type: str): + """Return (bucket, actual_key), falling back to a deprecated spelling.""" + if root.get(type): + return root[type], type + + for legacy in _BUCKET_ALIASES.get(type, ()): + if root.get(legacy): + if legacy not in _warned_aliases: + _warned_aliases.add(legacy) + logger.warning( + f"[Dataflex] config section '{legacy}:' is deprecated, rename it to '{type}:'. " + f"It still works for now." + ) + return root[legacy], legacy + + return {}, type + + def load_component(type: str, cfg_file: str, name: str, runtime_vars: Optional[Dict[str, str]] = None) -> Dict[str, Any]: with open(cfg_file, "r", encoding="utf-8") as f: root = yaml.safe_load(f) or {} - bucket = (root.get(type) or {}) + bucket, key = _resolve_bucket(root, type) if name not in bucket: available = ", ".join(sorted(bucket.keys())) - raise ValueError(f"{type} '{name}' not found. Available: {available}") + raise ValueError(f"{key} '{name}' not found. Available: {available}") params = dict(bucket[name].get("params") or {}) - # 简单占位替换(如 ${output_dir}) + # Simple placeholder substitution (e.g. ${output_dir}) if runtime_vars: def subst(v): if isinstance(v, str): diff --git a/src/dataflex/utils/selector_io.py b/src/dataflex/utils/selector_io.py index 270c20a..4f4ea06 100644 --- a/src/dataflex/utils/selector_io.py +++ b/src/dataflex/utils/selector_io.py @@ -29,8 +29,8 @@ def save_selection( accelerator, ) -> None: """ - 以统一格式保存,并仅由主进程落盘。 - 存储为标准的 JSON 格式。 + Save in a unified format and only by the main process. + Stored as standard JSON format. """ if accelerator.is_main_process: _ensure_parent_dir(save_path) @@ -39,5 +39,5 @@ def save_selection( "metric": metric, } with open(save_path, "w") as f: - json.dump(payload, f, indent=4) # 保存为漂亮的JSON格式 + json.dump(payload, f, indent=4) # Save in a pretty JSON format logger.info(f"[Dataflex] Saved selection to {save_path}.")