Update CTransformersComponent build_config() method
This commit is contained in:
parent
239e4bcb7c
commit
306292c36f
1 changed files with 5 additions and 35 deletions
|
|
@ -11,46 +11,16 @@ class CTransformersComponent(CustomComponent):
|
||||||
def build_config(self):
|
def build_config(self):
|
||||||
return {
|
return {
|
||||||
"model": {"display_name": "Model", "required": True},
|
"model": {"display_name": "Model", "required": True},
|
||||||
"model_file": {"display_name": "Model File", "required": False},
|
"model_file": {"display_name": "Model File", "required": False,"field_type":"file", "file_types":[".bin"]},
|
||||||
"model_type": {"display_name": "Model Type", "required": False},
|
"model_type": {"display_name": "Model Type", "required": True},
|
||||||
"config": {"display_name": "Config", "advanced": True, "required": False,"field_type":"dict","value":'{"top_k":40,"top_p":0.95,"temperature":0.8,"repetition_penalty":1.1,"last_n_tokens":64,"seed":-1,"max_new_tokens":256,"stop":"","stream":"False","reset":"True","batch_size":8,"threads":-1,"context_length":-1,"gpu_layers":0}'}
|
"config": {"display_name": "Config", "advanced": True, "required": False,"field_type":"dict","value":'{"top_k":40,"top_p":0.95,"temperature":0.8,"repetition_penalty":1.1,"last_n_tokens":64,"seed":-1,"max_new_tokens":256,"stop":"","stream":"False","reset":"True","batch_size":8,"threads":-1,"context_length":-1,"gpu_layers":0}'}
|
||||||
}
|
}
|
||||||
|
|
||||||
def build(
|
def build(
|
||||||
self,
|
self,
|
||||||
model: str,
|
model: str,
|
||||||
model_file: Optional[str] = None,
|
model_file: str,
|
||||||
model_type: Optional[str] = None,
|
model_type: str,
|
||||||
config: Optional[Dict] = None
|
config: Optional[Dict] = None
|
||||||
) -> CTransformers:
|
) -> CTransformers:
|
||||||
# Default config values
|
return CTransformers(model=model, model_file=model_file, model_type=model_type, config=config)
|
||||||
default_config = {
|
|
||||||
"top_k": 40,
|
|
||||||
"top_p": 0.95,
|
|
||||||
"temperature": 0.8,
|
|
||||||
"repetition_penalty": 1.1,
|
|
||||||
"last_n_tokens": 64,
|
|
||||||
"seed": -1,
|
|
||||||
"max_new_tokens": 256,
|
|
||||||
"stop": None,
|
|
||||||
"stream": False,
|
|
||||||
"reset": True,
|
|
||||||
"batch_size": 8,
|
|
||||||
"threads": -1,
|
|
||||||
"context_length": -1,
|
|
||||||
"gpu_layers": 0
|
|
||||||
}
|
|
||||||
|
|
||||||
# If there is a custom config, update the default config with it
|
|
||||||
if config:
|
|
||||||
default_config.update(config)
|
|
||||||
|
|
||||||
# Assuming the import below is correct and CTransformers is a class within the langchain library
|
|
||||||
# that inherits from BaseLanguageModel. The following import statement is required:
|
|
||||||
# from langchain.llms.integration_module import CTransformers
|
|
||||||
|
|
||||||
return CTransformers(model=model, model_file=model_file, model_type=model_type, config=default_config)
|
|
||||||
|
|
||||||
# Note: The actual CTransformers class needs to be imported from the correct module inside the langchain library.
|
|
||||||
# The `integration_module` in the import statement is just a placeholder and should be replaced with
|
|
||||||
# the actual module where the CTransformers class is located.
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue