mirror of
https://github.com/KoboldAI/KoboldAI-Client.git
synced 2025-04-26 15:58:47 +02:00
Fix no attribute get_checkpoint_shard_files
This commit is contained in:
parent
6e82f205b4
commit
d5ab3ef5b1
@ -1180,6 +1180,7 @@ if(not vars.use_colab_tpu and vars.model not in ["InferKit", "Colab", "OAI", "Go
|
|||||||
utils.aria2_hook(pretrained_model_name_or_path, **kwargs)
|
utils.aria2_hook(pretrained_model_name_or_path, **kwargs)
|
||||||
return old_from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs)
|
return old_from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs)
|
||||||
PreTrainedModel.from_pretrained = new_from_pretrained
|
PreTrainedModel.from_pretrained = new_from_pretrained
|
||||||
|
if(hasattr(modeling_utils, "get_checkpoint_shard_files")):
|
||||||
old_get_checkpoint_shard_files = modeling_utils.get_checkpoint_shard_files
|
old_get_checkpoint_shard_files = modeling_utils.get_checkpoint_shard_files
|
||||||
def new_get_checkpoint_shard_files(pretrained_model_name_or_path, index_filename, *args, **kwargs):
|
def new_get_checkpoint_shard_files(pretrained_model_name_or_path, index_filename, *args, **kwargs):
|
||||||
utils.num_shards = utils.get_num_shards(index_filename)
|
utils.num_shards = utils.get_num_shards(index_filename)
|
||||||
@ -1707,6 +1708,7 @@ else:
|
|||||||
utils.aria2_hook(pretrained_model_name_or_path, **kwargs)
|
utils.aria2_hook(pretrained_model_name_or_path, **kwargs)
|
||||||
return old_from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs)
|
return old_from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs)
|
||||||
PreTrainedModel.from_pretrained = new_from_pretrained
|
PreTrainedModel.from_pretrained = new_from_pretrained
|
||||||
|
if(hasattr(modeling_utils, "get_checkpoint_shard_files")):
|
||||||
old_get_checkpoint_shard_files = modeling_utils.get_checkpoint_shard_files
|
old_get_checkpoint_shard_files = modeling_utils.get_checkpoint_shard_files
|
||||||
def new_get_checkpoint_shard_files(pretrained_model_name_or_path, index_filename, *args, **kwargs):
|
def new_get_checkpoint_shard_files(pretrained_model_name_or_path, index_filename, *args, **kwargs):
|
||||||
utils.num_shards = utils.get_num_shards(index_filename)
|
utils.num_shards = utils.get_num_shards(index_filename)
|
||||||
|
Loading…
x
Reference in New Issue
Block a user