Integrate with Sentence Transformers v5.4

#9
by tomaarsen HF Staff - opened
Files changed (1) hide show
  1. sentence_transformers_impl.py +7 -3
sentence_transformers_impl.py CHANGED
@@ -55,17 +55,21 @@ class Transformer(nn.Module):
55
  if config_args is None:
56
  config_args = {}
57
 
 
 
 
 
 
58
  if not model_args.get("trust_remote_code", False):
59
  raise ValueError(
60
  "You need to set `trust_remote_code=True` to load this model."
61
  )
62
 
63
- self.config = AutoConfig.from_pretrained(model_name_or_path, **config_args, cache_dir=cache_dir)
64
- self.auto_model = AutoModel.from_pretrained(model_name_or_path, config=self.config, cache_dir=cache_dir, **model_args)
65
 
66
  self.tokenizer = AutoTokenizer.from_pretrained(
67
  model_name_or_path,
68
- cache_dir=cache_dir,
69
  **tokenizer_args,
70
  )
71
 
 
55
  if config_args is None:
56
  config_args = {}
57
 
58
+ if cache_dir is not None:
59
+ config_args["cache_dir"] = cache_dir
60
+ model_args["cache_dir"] = cache_dir
61
+ tokenizer_args["cache_dir"] = cache_dir
62
+
63
  if not model_args.get("trust_remote_code", False):
64
  raise ValueError(
65
  "You need to set `trust_remote_code=True` to load this model."
66
  )
67
 
68
+ self.config = AutoConfig.from_pretrained(model_name_or_path, **config_args)
69
+ self.auto_model = AutoModel.from_pretrained(model_name_or_path, config=self.config, **model_args)
70
 
71
  self.tokenizer = AutoTokenizer.from_pretrained(
72
  model_name_or_path,
 
73
  **tokenizer_args,
74
  )
75