📄
🤗 Transformers 過去PR再現調査: HfArgumentParser Union 型順序依存バグ (#39467)
PR 概要
TrainingArguments の一部フィールドで、コマンドラインから str で与えた JSON 文字列を HfArgumentParser が dict へ解釈しようとした際に失敗する不具合を調査・再現し、挙動と原因を整理したメモ。例えば、 --accelerator_config(型: Optional[Union[dict, str]] )で、Union 内の型処理ロジックが特定順序に依存しているため発生していた。
再現
- PR 直前のコミットに戻りブランチ作成
git checkout bda75b4011239d065de84aa3e744b67ebfa7b245
git switch -c pr_39462_hfargumentparser_local_path_str_error
以下コードを実行し、--accelerator_config に JSON 文字列を渡すとエラー。
from transformers import (
HfArgumentParser,
Trainer,
TrainingArguments,
)
def main():
parser = HfArgumentParser(TrainingArguments)
parser.parse_args_into_dataclasses()[0]
print("parse success")
if __name__ == "__main__":
main()
python3 playground.py --accelerator_config '{"gradient_accumulation_kwargs": {"num_steps": 2}}'
error: argument --accelerator_config/--accelerator-config: invalid dict value: '{"gradient_accumulation_kwargs": {"num_steps": 2}}'
原因分析
該当処理は hf_argparser.py の HfArgumentParser._parse_dataclass_field(185行目付近)。Union 型フィールドを扱う際に
-
Union[T, NoneType]を想定した分岐でNoneTypeを除去しようとしている - 実際には
Union[dict, str, NoneType]のように 3 要素以上でも同じロジックが走る - 実装は
field.type.__args__[0]とfield.type.__args__[1]のみを参照して判定するため、順序次第で本来保持すべきstrが脱落 - その結果、後段で「
dictとして解釈しようとして失敗 → invalid dict value」になる
問題となるコード(要旨)は次:
field.type = (
field.type.__args__[0] if isinstance(None, field.type.__args__[1]) else field.type.__args__[1]
)
Union が 2 要素(T | None)である前提に依存しており、3 要素以上(dict | str | None)を安全に正規化できていない点が根本。
なぜ PR の変更で直るか
PR ではロジックを直さず、ソースコード内の Union[str, dict] を Union[dict, str] に並べ替えている。これにより上記の簡易分岐で残される要素が str になり、入力文字列が弾かれなくなる(副次的解決)。
なお、なぜソースコード内の Union[str, dict] を Union[dict, str] に並べ替えると直るのかは、以下の記事を参照。
本質的な修正の方向(補足メモ)
- この PR 作成者が、別の PR のコメントで修正の方向性を示唆している。
- 現在の
HfArgumentParser._parse_dataclass_fieldはそもそもOptional[Union[dict, str]]型を上手く扱うことを想定しておらず、さらにUnion[dict, str]やUnion[str, dict]も上手く扱えない。そのため、それらを想定した修正が必要と考える。
関連リンク
- この PR の問題が起きたのは、
Union[str, typing.Dict]をUnion[str, dict]に書き換えたことが関連している。
Discussion