zodiac.streams.token_stream

 1#  # # <!-- // /*  SPDX-License-Identifier: MPL-2.0*/ -->
 2#  # # <!-- // /*  d a r k s h a p e s */ -->
 3
 4# pylint: disable=import-error
 5
 6
 7import warnings
 8from pathlib import Path
 9from typing import Callable, Optional
10
11warnings.filterwarnings("ignore", category=DeprecationWarning)
12
13from litellm.utils import create_tokenizer, token_counter
14from toga.sources import Source
15from zodiac.providers.registry_entry import RegistryEntry
16
17
18class TokenStream(Source):
19    def __init__(self):
20        self.tokenizer: Optional[str] = None
21        self.message: Optional[str] = None
22        self.tokenizer_args = {}
23
24    async def set_tokenizer(self, registry_entry: RegistryEntry) -> Callable:
25        """Pass message to model routine\n
26        :param model: Path to model
27        :param message: Text to encode
28        :return: Token embeddings for the model"""
29
30        import json
31
32        if registry_entry.tokenizer:
33            with open(str(registry_entry.tokenizer), encoding="UTF-8") as file_obj:
34                tokenizer_json = json.load(file_obj)
35                tokenizer_data = json.dumps(tokenizer_json)
36                self.tokenizer_args = {"custom_tokenizer": create_tokenizer(tokenizer_data)}
37        else:
38            # model_name = os.path.split(registry_entry.model)
39            # model_name = os.path.join(os.path.split(model_name[0])[-1], model_name[-1])
40            # self.status_log.registry_entry.model
41            self.tokenizer_args = {"model": registry_entry.model}
42
43    async def token_count(
44        self,
45        message: str,
46    ) -> Callable:
47        """Return token count of message based on model\n
48        :param model: Model path to lookup tokenizer for
49        :param message: Message to tokenize
50        :return: `int` Number of tokens needed to represent message"""
51        import warnings
52
53        warnings.filterwarnings("ignore", category=DeprecationWarning)
54        character_count = len(message)
55        return token_counter(text=message, **self.tokenizer_args), character_count
class TokenStream(toga.sources.base.Source):
19class TokenStream(Source):
20    def __init__(self):
21        self.tokenizer: Optional[str] = None
22        self.message: Optional[str] = None
23        self.tokenizer_args = {}
24
25    async def set_tokenizer(self, registry_entry: RegistryEntry) -> Callable:
26        """Pass message to model routine\n
27        :param model: Path to model
28        :param message: Text to encode
29        :return: Token embeddings for the model"""
30
31        import json
32
33        if registry_entry.tokenizer:
34            with open(str(registry_entry.tokenizer), encoding="UTF-8") as file_obj:
35                tokenizer_json = json.load(file_obj)
36                tokenizer_data = json.dumps(tokenizer_json)
37                self.tokenizer_args = {"custom_tokenizer": create_tokenizer(tokenizer_data)}
38        else:
39            # model_name = os.path.split(registry_entry.model)
40            # model_name = os.path.join(os.path.split(model_name[0])[-1], model_name[-1])
41            # self.status_log.registry_entry.model
42            self.tokenizer_args = {"model": registry_entry.model}
43
44    async def token_count(
45        self,
46        message: str,
47    ) -> Callable:
48        """Return token count of message based on model\n
49        :param model: Model path to lookup tokenizer for
50        :param message: Message to tokenize
51        :return: `int` Number of tokens needed to represent message"""
52        import warnings
53
54        warnings.filterwarnings("ignore", category=DeprecationWarning)
55        character_count = len(message)
56        return token_counter(text=message, **self.tokenizer_args), character_count

A base class for data sources, providing an implementation of data notifications.

tokenizer: Optional[str]
message: Optional[str]
tokenizer_args
async def set_tokenizer( self, registry_entry: zodiac.providers.registry_entry.RegistryEntry) -> Callable:
25    async def set_tokenizer(self, registry_entry: RegistryEntry) -> Callable:
26        """Pass message to model routine\n
27        :param model: Path to model
28        :param message: Text to encode
29        :return: Token embeddings for the model"""
30
31        import json
32
33        if registry_entry.tokenizer:
34            with open(str(registry_entry.tokenizer), encoding="UTF-8") as file_obj:
35                tokenizer_json = json.load(file_obj)
36                tokenizer_data = json.dumps(tokenizer_json)
37                self.tokenizer_args = {"custom_tokenizer": create_tokenizer(tokenizer_data)}
38        else:
39            # model_name = os.path.split(registry_entry.model)
40            # model_name = os.path.join(os.path.split(model_name[0])[-1], model_name[-1])
41            # self.status_log.registry_entry.model
42            self.tokenizer_args = {"model": registry_entry.model}

Pass message to model routine

Parameters
  • model: Path to model
  • message: Text to encode
Returns

Token embeddings for the model

async def token_count(self, message: str) -> Callable:
44    async def token_count(
45        self,
46        message: str,
47    ) -> Callable:
48        """Return token count of message based on model\n
49        :param model: Model path to lookup tokenizer for
50        :param message: Message to tokenize
51        :return: `int` Number of tokens needed to represent message"""
52        import warnings
53
54        warnings.filterwarnings("ignore", category=DeprecationWarning)
55        character_count = len(message)
56        return token_counter(text=message, **self.tokenizer_args), character_count

Return token count of message based on model

Parameters
  • model: Model path to lookup tokenizer for
  • message: Message to tokenize
Returns

int Number of tokens needed to represent message