From 8e61e2637020d515f57d3d58ec6f52aba43cfa2a Mon Sep 17 00:00:00 2001 From: Steven Liu Date: Fri, 31 Mar 2023 14:41:14 -0700 Subject: [PATCH] fix kwargs --- src/peft/peft_model.py | 2 +- src/peft/utils/config.py | 8 ++++---- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index 2b79f4c..0afd047 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -92,7 +92,7 @@ class PeftModel(PushToHubMixin, torch.nn.Module): save_directory (`str`): Directory where the adapter model and configuration files will be saved (will be created if it does not exist). - **kwargs: + kwargs (additional keyword arguments, *optional*): Additional keyword arguments passed along to the `push_to_hub` method. """ if os.path.isfile(save_directory): diff --git a/src/peft/utils/config.py b/src/peft/utils/config.py index 3ace67a..bdd7277 100644 --- a/src/peft/utils/config.py +++ b/src/peft/utils/config.py @@ -65,7 +65,7 @@ class PeftConfigMixin(PushToHubMixin): Args: save_directory (`str`): The directory where the configuration will be saved. - kwargs: + kwargs (additional keyword arguments, *optional*): Additional keyword arguments passed along to the [`~transformers.utils.PushToHubMixin.push_to_hub`] method. """ @@ -89,7 +89,7 @@ class PeftConfigMixin(PushToHubMixin): Args: pretrained_model_name_or_path (`str`): The directory or the Hub repository id where the configuration is saved. - kwargs: + kwargs (additional keyword arguments, *optional*): Additional keyword arguments passed along to the child class initialization. """ if os.path.isfile(os.path.join(pretrained_model_name_or_path, CONFIG_NAME)): @@ -145,8 +145,8 @@ class PeftConfig(PeftConfigMixin): @dataclass class PromptLearningConfig(PeftConfig): """ - This is the base configuration class to store the configuration of a Union[[`~peft.PrefixTuning`], - [`~peft.PromptEncoder`], [`~peft.PromptTuning`]]. + This is the base configuration class to store the configuration of [`PrefixTuning`], [`PromptEncoder`], or + [`PromptTuning`]. Args: num_virtual_tokens (`int`): The number of virtual tokens to use.