Move langauge model abstractions and implementation into their own module; create a new README; update dependencies
This commit is contained in:
@@ -1,75 +1,18 @@
|
||||
# Controller
|
||||
|
||||
The Controller class is responsible for traversing the Graph of Operations (GoO), which is a static structure that is constructed once, before the execution starts.
|
||||
The Controller class is responsible for traversing the Graph of Operations (GoO), which is a static structure that is constructed once, before the execution starts.
|
||||
GoO prescribes the execution plan of thought operations and the Controller invokes their execution, generating the Graph Reasoning State (GRS).
|
||||
|
||||
In order for a GoO to be executed, an instance of Large Language Model (LLM) must be supplied to the controller.
|
||||
Currently, the framework supports the following LLMs:
|
||||
- GPT-4 / GPT-3.5 (Remote - OpenAI API)
|
||||
- Llama-2 (Local - HuggingFace Transformers)
|
||||
In order for a GoO to be executed, an instance of Large Language Model (LLM) must be supplied to the controller (along with other required objects).
|
||||
Please refer to the [Language Models](../language_models/README.md) section for more information about LLMs.
|
||||
|
||||
The following section describes how to instantiate individual LLMs and the Controller to run a defined GoO.
|
||||
Furthermore, the process of adding new LLMs into the framework is outlined at the end.
|
||||
|
||||
## LLM Instantiation
|
||||
- Create a copy of `config_template.json` named `config.json`.
|
||||
- Fill configuration details based on the used model (below).
|
||||
|
||||
### GPT-4 / GPT-3.5
|
||||
- Adjust predefined `chatgpt`, `chatgpt4` or create new configuration with an unique key.
|
||||
|
||||
| Key | Value |
|
||||
|---------------------|---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| model_id | Model name based on [OpenAI model overview](https://platform.openai.com/docs/models/overview). |
|
||||
| prompt_token_cost | Price per 1000 prompt tokens based on [OpenAI pricing](https://openai.com/pricing), used for calculating cumulative price per LLM instance. |
|
||||
| response_token_cost | Price per 1000 response tokens based on [OpenAI pricing](https://openai.com/pricing), used for calculating cumulative price per LLM instance. |
|
||||
| temperature | Parameter of OpenAI models that controls randomness and the creativity of the responses (higher temperature = more diverse and unexpected responses). Value between 0.0 and 2.0, default is 1.0. More information can be found in the [OpenAI API reference](https://platform.openai.com/docs/api-reference/completions/create#completions/create-temperature). |
|
||||
| max_tokens | The maximum number of tokens to generate in the chat completion. Value depends on the maximum context size of the model specified in the [OpenAI model overview](https://platform.openai.com/docs/models/overview). More information can be found in the [OpenAI API reference](https://platform.openai.com/docs/api-reference/chat/create#chat/create-max_tokens). |
|
||||
| stop | String or array of strings specifying sequence of characters which if detected, stops further generation of tokens. More information can be found in the [OpenAI API reference](https://platform.openai.com/docs/api-reference/chat/create#chat/create-stop). |
|
||||
| organization | Organization to use for the API requests (may be empty). |
|
||||
| api_key | Personal API key that will be used to access OpenAI API. |
|
||||
|
||||
- Instantiate the language model based on the selected configuration key (predefined / custom).
|
||||
```
|
||||
lm = controller.ChatGPT(
|
||||
"path/to/config.json",
|
||||
model_name=<configuration key>
|
||||
)
|
||||
```
|
||||
|
||||
### Llama-2
|
||||
- Requires local hardware to run inference and a HuggingFace account.
|
||||
- Adjust predefined `llama7b-hf`, `llama13b-hf`, `llama70b-hf` or create a new configuration with an unique key.
|
||||
|
||||
| Key | Value |
|
||||
|---------------------|---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| model_id | Specifies HuggingFace Llama 2 model identifier (`meta-llama/<model_id>`). |
|
||||
| cache_dir | Local directory where model will be downloaded and accessed. |
|
||||
| prompt_token_cost | Price per 1000 prompt tokens (currently not used - local model = no cost). |
|
||||
| response_token_cost | Price per 1000 response tokens (currently not used - local model = no cost). |
|
||||
| temperature | Parameter that controls randomness and the creativity of the responses (higher temperature = more diverse and unexpected responses). Value between 0.0 and 1.0, default is 0.6. |
|
||||
| top_k | Top-K sampling method described in [Transformers tutorial](https://huggingface.co/blog/how-to-generate). Default value is set to 10. |
|
||||
| max_tokens | The maximum number of tokens to generate in the chat completion. More tokens require more memory. |
|
||||
|
||||
- Instantiate the language model based on the selected configuration key (predefined / custom).
|
||||
```
|
||||
lm = controller.Llama2HF(
|
||||
"path/to/config.json",
|
||||
model_name=<configuration key>
|
||||
)
|
||||
```
|
||||
- Request access to Llama-2 via the [Meta form](https://ai.meta.com/resources/models-and-libraries/llama-downloads/) using the same email address as for the HuggingFace account.
|
||||
- After the access is granted, go to [HuggingFace Llama-2 model card](https://huggingface.co/meta-llama/Llama-2-7b-chat-hf), log in and accept the license (_"You have been granted access to this model"_ message should appear).
|
||||
- Generate HuggingFace access token.
|
||||
- Log in from CLI with: `huggingface-cli login --token <your token>`.
|
||||
|
||||
Note: 4-bit quantization is used to reduce the model size for inference. During instantiation, the model is downloaded from HuggingFace into the cache directory specified in the `config.json`. Running queries using larger models will require multiple GPUs (splitting across many GPUs is done automatically by the Transformers library).
|
||||
The following section describes how to instantiate the Controller to run a defined GoO.
|
||||
|
||||
## Controller Instantiation
|
||||
- Requires custom `Prompter`, `Parser` and instantiated `GraphOfOperations` - creation of these is described separately.
|
||||
- Use instantiated `lm` from above.
|
||||
- Requires custom `Prompter`, `Parser`, as well as instantiated `GraphOfOperations` and `AbstractLanguageModel` - creation of these is described separately.
|
||||
- Prepare initial state (thought) as dictionary - this can be used in the initial prompts by the operations.
|
||||
```
|
||||
lm = ...create
|
||||
graph_of_operations = ...create
|
||||
|
||||
executor = controller.Controller(
|
||||
@@ -83,35 +26,3 @@ executor.run()
|
||||
executor.output_graph("path/to/output.json")
|
||||
```
|
||||
- After the run the graph is written to an output file, which contains individual operations, their thoughts, information about scores and validity and total amount of used tokens / cost.
|
||||
|
||||
## Adding LLMs
|
||||
More LLMs can be added by following these steps:
|
||||
- Create new class as a subclass of `AbstractLanguageModel`.
|
||||
- Use the constructor for loading configuration and instantiating the language model (if needed).
|
||||
```
|
||||
class CustomLanguageModel(AbstractLanguageModel):
|
||||
def __init__(
|
||||
self,
|
||||
config_path: str = "",
|
||||
model_name: str = "llama7b-hf",
|
||||
cache: bool = False
|
||||
) -> None:
|
||||
super().__init__(config_path, model_name, cache)
|
||||
self.config: Dict = self.config[model_name]
|
||||
|
||||
# Load data from configuration into variables if needed
|
||||
|
||||
# Instantiate LLM if needed
|
||||
```
|
||||
- Implement `query` abstract method that is used to get a list of responses from the LLM (call to remote API or local model inference).
|
||||
```
|
||||
def query(self, query: str, num_responses: int = 1) -> Any:
|
||||
# Support caching
|
||||
# Call LLM and retrieve list of responses - based on num_responses
|
||||
# Return LLM response structure (not only raw strings)
|
||||
```
|
||||
- Implement `get_response_texts` abstract method that is used to get a list of raw texts from the LLM response structure produced by `query`.
|
||||
```
|
||||
def get_response_texts(self, query_response: Union[List[Dict], Dict]) -> List[str]:
|
||||
# Retrieve list of raw strings from the LLM response structure
|
||||
```
|
||||
|
||||
@@ -1,4 +1 @@
|
||||
from .chatgpt import ChatGPT
|
||||
from .llamachat_hf import Llama2HF
|
||||
from .abstract_language_model import AbstractLanguageModel
|
||||
from .controller import Controller
|
||||
|
||||
@@ -1,92 +0,0 @@
|
||||
# Copyright (c) 2023 ETH Zurich.
|
||||
# All rights reserved.
|
||||
#
|
||||
# Use of this source code is governed by a BSD-style license that can be
|
||||
# found in the LICENSE file.
|
||||
#
|
||||
# main author: Nils Blach
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Dict, Union, Any
|
||||
import json
|
||||
import os
|
||||
import logging
|
||||
|
||||
|
||||
class AbstractLanguageModel(ABC):
|
||||
"""
|
||||
Abstract base class that defines the interface for all language models.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, config_path: str = "", model_name: str = "", cache: bool = False
|
||||
) -> None:
|
||||
"""
|
||||
Initialize the AbstractLanguageModel instance with configuration, model details, and caching options.
|
||||
|
||||
:param config_path: Path to the config file. Defaults to "".
|
||||
:type config_path: str
|
||||
:param model_name: Name of the language model. Defaults to "".
|
||||
:type model_name: str
|
||||
:param cache: Flag to determine whether to cache responses. Defaults to False.
|
||||
:type cache: bool
|
||||
"""
|
||||
self.logger = logging.getLogger(self.__class__.__name__)
|
||||
self.config: Dict = None
|
||||
self.model_name: str = model_name
|
||||
self.cache = cache
|
||||
if self.cache:
|
||||
self.respone_cache: Dict[str, List[Any]] = {}
|
||||
self.load_config(config_path)
|
||||
self.prompt_tokens: int = 0
|
||||
self.completion_tokens: int = 0
|
||||
self.cost: float = 0.0
|
||||
|
||||
def load_config(self, path: str) -> None:
|
||||
"""
|
||||
Load configuration from a specified path.
|
||||
|
||||
:param path: Path to the config file. If an empty path provided,
|
||||
default is `config.json` in the current directory.
|
||||
:type path: str
|
||||
"""
|
||||
if path == "":
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
path = os.path.join(current_dir, "config.json")
|
||||
|
||||
with open(path, "r") as f:
|
||||
self.config = json.load(f)
|
||||
|
||||
self.logger.debug(f"Loaded config from {path} for {self.model_name}")
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
"""
|
||||
Clear the response cache.
|
||||
"""
|
||||
self.respone_cache.clear()
|
||||
|
||||
@abstractmethod
|
||||
def query(self, query: str, num_responses: int = 1) -> Any:
|
||||
"""
|
||||
Abstract method to query the language model.
|
||||
|
||||
:param query: The query to be posed to the language model.
|
||||
:type query: str
|
||||
:param num_responses: The number of desired responses.
|
||||
:type num_responses: int
|
||||
:return: The language model's response(s).
|
||||
:rtype: Any
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_response_texts(self, query_responses: Union[List[Dict], Dict]) -> List[str]:
|
||||
"""
|
||||
Abstract method to extract response texts from the language model's response(s).
|
||||
|
||||
:param query_responses: The responses returned from the language model.
|
||||
:type query_responses: Union[List[Dict], Dict]
|
||||
:return: List of textual responses.
|
||||
:rtype: List[str]
|
||||
"""
|
||||
pass
|
||||
@@ -1,156 +0,0 @@
|
||||
# Copyright (c) 2023 ETH Zurich.
|
||||
# All rights reserved.
|
||||
#
|
||||
# Use of this source code is governed by a BSD-style license that can be
|
||||
# found in the LICENSE file.
|
||||
#
|
||||
# main author: Nils Blach
|
||||
|
||||
import backoff
|
||||
import openai
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from typing import List, Dict, Union
|
||||
|
||||
from .abstract_language_model import AbstractLanguageModel
|
||||
|
||||
|
||||
class ChatGPT(AbstractLanguageModel):
|
||||
"""
|
||||
The ChatGPT class handles interactions with the OpenAI models using the provided configuration.
|
||||
|
||||
Inherits from the AbstractLanguageModel and implements its abstract methods.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, config_path: str = "", model_name: str = "chatgpt", cache: bool = False
|
||||
) -> None:
|
||||
"""
|
||||
Initialize the ChatGPT instance with configuration, model details, and caching options.
|
||||
|
||||
:param config_path: Path to the configuration file. Defaults to "".
|
||||
:type config_path: str
|
||||
:param model_name: Name of the model, default is 'chatgpt'. Used to select the correct configuration.
|
||||
:type model_name: str
|
||||
:param cache: Flag to determine whether to cache responses. Defaults to False.
|
||||
:type cache: bool
|
||||
"""
|
||||
super().__init__(config_path, model_name, cache)
|
||||
self.config: Dict = self.config[model_name]
|
||||
# The model_id is the id of the model that is used for chatgpt, i.e. gpt-4, gpt-3.5-turbo, etc.
|
||||
self.model_id: str = self.config["model_id"]
|
||||
# The prompt_token_cost and response_token_cost are the costs for 1000 prompt tokens and 1000 response tokens respectively.
|
||||
self.prompt_token_cost: float = self.config["prompt_token_cost"]
|
||||
self.response_token_cost: float = self.config["response_token_cost"]
|
||||
# The temperature of a model is defined as the randomness of the model's output.
|
||||
self.temperature: float = self.config["temperature"]
|
||||
# The maximum number of tokens to generate in the chat completion.
|
||||
self.max_tokens: int = self.config["max_tokens"]
|
||||
# The stop sequence is a sequence of tokens that the model will stop generating at (it will not generate the stop sequence).
|
||||
self.stop: Union[str, List[str]] = self.config["stop"]
|
||||
# The account organization is the organization that is used for chatgpt.
|
||||
self.organization: str = self.config["organization"]
|
||||
if self.organization == "":
|
||||
self.logger.warning("OPENAI_ORGANIZATION is not set")
|
||||
else:
|
||||
openai.organization = self.organization
|
||||
# The api key is the api key that is used for chatgpt. Env variable OPENAI_API_KEY takes precedence over config.
|
||||
self.api_key: str = os.getenv("OPENAI_API_KEY", self.config["api_key"])
|
||||
if self.api_key == "":
|
||||
raise ValueError("OPENAI_API_KEY is not set")
|
||||
openai.api_key = self.api_key
|
||||
|
||||
def query(self, query: str, num_responses: int = 1) -> Dict:
|
||||
"""
|
||||
Query the OpenAI model for responses.
|
||||
|
||||
:param query: The query to be posed to the language model.
|
||||
:type query: str
|
||||
:param num_responses: Number of desired responses, default is 1.
|
||||
:type num_responses: int
|
||||
:return: Response(s) from the OpenAI model.
|
||||
:rtype: Dict
|
||||
"""
|
||||
if self.cache and query in self.respone_cache:
|
||||
return self.respone_cache[query]
|
||||
|
||||
if num_responses == 1:
|
||||
response = self.chat([{"role": "user", "content": query}], num_responses)
|
||||
else:
|
||||
response = []
|
||||
next_try = num_responses
|
||||
total_num_attempts = num_responses
|
||||
while num_responses > 0 and total_num_attempts > 0:
|
||||
try:
|
||||
assert next_try > 0
|
||||
res = self.chat([{"role": "user", "content": query}], next_try)
|
||||
response.append(res)
|
||||
num_responses -= next_try
|
||||
next_try = min(num_responses, next_try)
|
||||
except Exception as e:
|
||||
next_try = (next_try + 1) // 2
|
||||
self.logger.warning(
|
||||
f"Error in chatgpt: {e}, trying again with {next_try} samples"
|
||||
)
|
||||
time.sleep(random.randint(1, 3))
|
||||
total_num_attempts -= 1
|
||||
|
||||
if self.cache:
|
||||
self.respone_cache[query] = response
|
||||
return response
|
||||
|
||||
@backoff.on_exception(
|
||||
backoff.expo, openai.error.OpenAIError, max_time=10, max_tries=6
|
||||
)
|
||||
def chat(self, messages: List[Dict], num_responses: int = 1) -> Dict:
|
||||
"""
|
||||
Send chat messages to the OpenAI model and retrieves the model's response.
|
||||
Implements backoff on OpenAI error.
|
||||
|
||||
:param messages: A list of message dictionaries for the chat.
|
||||
:type messages: List[Dict]
|
||||
:param num_responses: Number of desired responses, default is 1.
|
||||
:type num_responses: int
|
||||
:return: The OpenAI model's response.
|
||||
:rtype: Dict
|
||||
"""
|
||||
response = openai.ChatCompletion.create(
|
||||
model=self.model_id,
|
||||
messages=messages,
|
||||
temperature=self.temperature,
|
||||
max_tokens=self.max_tokens,
|
||||
n=num_responses,
|
||||
stop=self.stop,
|
||||
)
|
||||
|
||||
self.prompt_tokens += response["usage"]["prompt_tokens"]
|
||||
self.completion_tokens += response["usage"]["completion_tokens"]
|
||||
prompt_tokens_k = float(self.prompt_tokens) / 1000.0
|
||||
completion_tokens_k = float(self.completion_tokens) / 1000.0
|
||||
self.cost = (
|
||||
self.prompt_token_cost * prompt_tokens_k
|
||||
+ self.response_token_cost * completion_tokens_k
|
||||
)
|
||||
self.logger.info(
|
||||
f"This is the response from chatgpt: {response}"
|
||||
f"\nThis is the cost of the response: {self.cost}"
|
||||
)
|
||||
return response
|
||||
|
||||
def get_response_texts(self, query_response: Union[List[Dict], Dict]) -> List[str]:
|
||||
"""
|
||||
Extract the response texts from the query response.
|
||||
|
||||
:param query_response: The response dictionary (or list of dictionaries) from the OpenAI model.
|
||||
:type query_response: Union[List[Dict], Dict]
|
||||
:return: List of response strings.
|
||||
:rtype: List[str]
|
||||
"""
|
||||
if isinstance(query_response, Dict):
|
||||
query_response = [query_response]
|
||||
return [
|
||||
choice["message"]["content"]
|
||||
for response in query_response
|
||||
for choice in response["choices"]
|
||||
]
|
||||
@@ -1,49 +0,0 @@
|
||||
{
|
||||
"chatgpt" : {
|
||||
"model_id": "gpt-3.5-turbo",
|
||||
"prompt_token_cost": 0.0015,
|
||||
"response_token_cost": 0.002,
|
||||
"temperature": 1.0,
|
||||
"max_tokens": 1536,
|
||||
"stop": null,
|
||||
"organization": "",
|
||||
"api_key": ""
|
||||
},
|
||||
"chatgpt4" : {
|
||||
"model_id": "gpt-4",
|
||||
"prompt_token_cost": 0.03,
|
||||
"response_token_cost": 0.06,
|
||||
"temperature": 1.0,
|
||||
"max_tokens": 4096,
|
||||
"stop": null,
|
||||
"organization": "",
|
||||
"api_key": ""
|
||||
},
|
||||
"llama7b-hf" : {
|
||||
"model_id": "Llama-2-7b-chat-hf",
|
||||
"cache_dir": "/llama",
|
||||
"prompt_token_cost": 0.0,
|
||||
"response_token_cost": 0.0,
|
||||
"temperature": 0.6,
|
||||
"top_k": 10,
|
||||
"max_tokens": 4096
|
||||
},
|
||||
"llama13b-hf" : {
|
||||
"model_id": "Llama-2-13b-chat-hf",
|
||||
"cache_dir": "/llama",
|
||||
"prompt_token_cost": 0.0,
|
||||
"response_token_cost": 0.0,
|
||||
"temperature": 0.6,
|
||||
"top_k": 10,
|
||||
"max_tokens": 4096
|
||||
},
|
||||
"llama70b-hf" : {
|
||||
"model_id": "Llama-2-70b-chat-hf",
|
||||
"cache_dir": "/llama",
|
||||
"prompt_token_cost": 0.0,
|
||||
"response_token_cost": 0.0,
|
||||
"temperature": 0.6,
|
||||
"top_k": 10,
|
||||
"max_tokens": 4096
|
||||
}
|
||||
}
|
||||
@@ -9,7 +9,7 @@
|
||||
import json
|
||||
import logging
|
||||
from typing import List
|
||||
from .abstract_language_model import AbstractLanguageModel
|
||||
from graph_of_thoughts.language_models import AbstractLanguageModel
|
||||
from graph_of_thoughts.operations import GraphOfOperations, Thought
|
||||
from graph_of_thoughts.prompter import Prompter
|
||||
from graph_of_thoughts.parser import Parser
|
||||
|
||||
@@ -1,119 +0,0 @@
|
||||
# Copyright (c) 2023 ETH Zurich.
|
||||
# All rights reserved.
|
||||
#
|
||||
# Use of this source code is governed by a BSD-style license that can be
|
||||
# found in the LICENSE file.
|
||||
#
|
||||
# main author: Ales Kubicek
|
||||
|
||||
import os
|
||||
import torch
|
||||
from typing import List, Dict, Union
|
||||
from .abstract_language_model import AbstractLanguageModel
|
||||
|
||||
|
||||
class Llama2HF(AbstractLanguageModel):
|
||||
"""
|
||||
An interface to use LLaMA 2 models through the HuggingFace library.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, config_path: str = "", model_name: str = "llama7b-hf", cache: bool = False
|
||||
) -> None:
|
||||
"""
|
||||
Initialize an instance of the Llama2HF class with configuration, model details, and caching options.
|
||||
|
||||
:param config_path: Path to the configuration file. Defaults to an empty string.
|
||||
:type config_path: str
|
||||
:param model_name: Specifies the name of the LLaMA model variant. Defaults to "llama7b-hf".
|
||||
Used to select the correct configuration.
|
||||
:type model_name: str
|
||||
:param cache: Flag to determine whether to cache responses. Defaults to False.
|
||||
:type cache: bool
|
||||
"""
|
||||
super().__init__(config_path, model_name, cache)
|
||||
self.config: Dict = self.config[model_name]
|
||||
# Detailed id of the used model.
|
||||
self.model_id: str = self.config["model_id"]
|
||||
# Costs for 1000 tokens.
|
||||
self.prompt_token_cost: float = self.config["prompt_token_cost"]
|
||||
self.response_token_cost: float = self.config["response_token_cost"]
|
||||
# The temperature is defined as the randomness of the model's output.
|
||||
self.temperature: float = self.config["temperature"]
|
||||
# Top K sampling.
|
||||
self.top_k: int = self.config["top_k"]
|
||||
# The maximum number of tokens to generate in the chat completion.
|
||||
self.max_tokens: int = self.config["max_tokens"]
|
||||
|
||||
# Important: must be done before importing transformers
|
||||
os.environ["TRANSFORMERS_CACHE"] = self.config["cache_dir"]
|
||||
import transformers
|
||||
|
||||
hf_model_id = f"meta-llama/{self.model_id}"
|
||||
model_config = transformers.AutoConfig.from_pretrained(hf_model_id)
|
||||
bnb_config = transformers.BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_quant_type="nf4",
|
||||
bnb_4bit_use_double_quant=True,
|
||||
bnb_4bit_compute_dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
self.tokenizer = transformers.AutoTokenizer.from_pretrained(hf_model_id)
|
||||
self.model = transformers.AutoModelForCausalLM.from_pretrained(
|
||||
hf_model_id,
|
||||
trust_remote_code=True,
|
||||
config=model_config,
|
||||
quantization_config=bnb_config,
|
||||
device_map="auto",
|
||||
)
|
||||
self.model.eval()
|
||||
torch.no_grad()
|
||||
|
||||
self.generate_text = transformers.pipeline(
|
||||
model=self.model, tokenizer=self.tokenizer, task="text-generation"
|
||||
)
|
||||
|
||||
def query(self, query: str, num_responses: int = 1) -> List[Dict]:
|
||||
"""
|
||||
Query the LLaMA 2 model for responses.
|
||||
|
||||
:param query: The query to be posed to the language model.
|
||||
:type query: str
|
||||
:param num_responses: Number of desired responses, default is 1.
|
||||
:type num_responses: int
|
||||
:return: Response(s) from the LLaMA 2 model.
|
||||
:rtype: List[Dict]
|
||||
"""
|
||||
if self.cache and query in self.respone_cache:
|
||||
return self.respone_cache[query]
|
||||
sequences = []
|
||||
query = f"<s><<SYS>>You are a helpful assistant. Always follow the intstructions precisely and output the response exactly in the requested format.<</SYS>>\n\n[INST] {query} [/INST]"
|
||||
for _ in range(num_responses):
|
||||
sequences.extend(
|
||||
self.generate_text(
|
||||
query,
|
||||
do_sample=True,
|
||||
top_k=self.top_k,
|
||||
num_return_sequences=1,
|
||||
eos_token_id=self.tokenizer.eos_token_id,
|
||||
max_length=self.max_tokens,
|
||||
)
|
||||
)
|
||||
response = [
|
||||
{"generated_text": sequence["generated_text"][len(query) :].strip()}
|
||||
for sequence in sequences
|
||||
]
|
||||
if self.cache:
|
||||
self.respone_cache[query] = response
|
||||
return response
|
||||
|
||||
def get_response_texts(self, query_responses: List[Dict]) -> List[str]:
|
||||
"""
|
||||
Extract the response texts from the query response.
|
||||
|
||||
:param query_responses: The response list of dictionaries generated from the `query` method.
|
||||
:type query_responses: List[Dict]
|
||||
:return: List of response strings.
|
||||
:rtype: List[str]
|
||||
"""
|
||||
return [query_response["generated_text"] for query_response in query_responses]
|
||||
Reference in New Issue
Block a user