From 37a66322593c6eaf742124e2bd6f6ec9bad21473 Mon Sep 17 00:00:00 2001 From: filipstrand Date: Wed, 14 Aug 2024 11:01:57 +0200 Subject: [PATCH] Return mlx arrays directly from tokenizer call --- src/flux_1_schnell/tokenizer/clip_tokenizer.py | 5 ++--- src/flux_1_schnell/tokenizer/t5_tokenizer.py | 5 ++--- 2 files changed, 4 insertions(+), 6 deletions(-) diff --git a/src/flux_1_schnell/tokenizer/clip_tokenizer.py b/src/flux_1_schnell/tokenizer/clip_tokenizer.py index 10f3918..2508241 100644 --- a/src/flux_1_schnell/tokenizer/clip_tokenizer.py +++ b/src/flux_1_schnell/tokenizer/clip_tokenizer.py @@ -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()) diff --git a/src/flux_1_schnell/tokenizer/t5_tokenizer.py b/src/flux_1_schnell/tokenizer/t5_tokenizer.py index 15ccbff..e227505 100644 --- a/src/flux_1_schnell/tokenizer/t5_tokenizer.py +++ b/src/flux_1_schnell/tokenizer/t5_tokenizer.py @@ -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())