diff --git a/kubeflow/trainer/backends/kubernetes/utils.py b/kubeflow/trainer/backends/kubernetes/utils.py index 66fdc481e..2f1ea93cc 100644 --- a/kubeflow/trainer/backends/kubernetes/utils.py +++ b/kubeflow/trainer/backends/kubernetes/utils.py @@ -492,6 +492,8 @@ def get_args_using_torchtune_config( ) -> list[str]: """ Get the Trainer args from the TorchTuneConfig. + This function generates arguments for TorchTune. It supports HuggingFace, + S3 and DataCache initializers to properly set the dataset arguments. """ args = [] @@ -532,6 +534,16 @@ def get_args_using_torchtune_config( else: args.append(f"dataset.data_dir={os.path.join(constants.DATASET_PATH, relative_path)}") + if isinstance(initializer, types.Initializer) and isinstance( + initializer.dataset, types.S3DatasetInitializer + ): + args.append(f"dataset.data_dir={initializer.dataset.storage_uri}") + + if isinstance(initializer, types.Initializer) and isinstance( + initializer.dataset, types.DataCacheInitializer + ): + args.append(f"dataset.data_dir={initializer.dataset.storage_uri}") + if fine_tuning_config.peft_config: args += get_args_from_peft_config(fine_tuning_config.peft_config)