Return mlx arrays directly from tokenizer call
This commit is contained in:
parent
5023437e89
commit
37a6632259
@ -9,13 +9,12 @@ class TokenizerCLIP:
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def tokenize(self, prompt: str) -> mx.array:
|
||||
text_input_ids = self.tokenizer(
|
||||
return self.tokenizer(
|
||||
[prompt],
|
||||
padding="max_length",
|
||||
max_length=TokenizerCLIP.MAX_TOKEN_LENGTH,
|
||||
truncation=True,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_tensors="pt",
|
||||
return_tensors="mlx",
|
||||
).input_ids
|
||||
return mx.array(text_input_ids.cpu().numpy())
|
||||
|
||||
@ -9,13 +9,12 @@ class TokenizerT5:
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def tokenize(self, prompt: str) -> mx.array:
|
||||
text_input_ids = self.tokenizer(
|
||||
return self.tokenizer(
|
||||
[prompt],
|
||||
padding="max_length",
|
||||
max_length=TokenizerT5.MAX_TOKEN_LENGTH,
|
||||
truncation=True,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_tensors="pt",
|
||||
return_tensors="mlx",
|
||||
).input_ids
|
||||
return mx.array(text_input_ids.cpu().numpy())
|
||||
|
||||
Loading…
Reference in New Issue
Block a user