@@ -324,23 +324,31 @@ def __init__(self, eos_token_id):
324324 self .eos_token = "</s>" if eos_token_id is not None else None
325325 self .eos_token_id = eos_token_id
326326
327- def __call__ (self , text , return_tensors = None , add_special_tokens = False ,):
327+ def __call__ (
328+ self ,
329+ text ,
330+ return_tensors = None ,
331+ add_special_tokens = False ,
332+ ):
328333 token_ids = list (range (len (text .split ())))
329334 if return_tensors == "pt" :
330335 return {"input_ids" : [token_ids ]}
331336 return {"input_ids" : token_ids }
332337
333- def decode (self , token_ids , skip_special_tokens = False ,):
338+ def decode (
339+ self ,
340+ token_ids ,
341+ skip_special_tokens = False ,
342+ ):
334343 return " " .join (f"word_{ i } " for i in token_ids )
335344
336345 for eos_token_id in (2 , None ):
337- loader = RawTextDataLoader (
338- MockTokenizer (eos_token_id ), chunk_size = 2048 , stride = 512
339- )
346+ loader = RawTextDataLoader (MockTokenizer (eos_token_id ), chunk_size = 2048 , stride = 512 )
340347 for text in ("" , " \n \t " ):
341- assert loader .smart_chunk_text (
342- text , chunk_size = 2048 , stride = 512 , return_tokenized = True
343- ) == [], f"empty input should yield no chunks (eos={ eos_token_id } , text={ text !r} )"
348+ assert (
349+ loader .smart_chunk_text (text , chunk_size = 2048 , stride = 512 , return_tokenized = True )
350+ == []
351+ ), f"empty input should yield no chunks (eos={ eos_token_id } , text={ text !r} )"
344352 assert loader .chunk_text (text ) == [], (
345353 f"chunk_text should yield no chunks for empty input "
346354 f"(eos={ eos_token_id } , text={ text !r} )"
0 commit comments