ludwig-ai--ludwig
593b94c120
pytest / Unit Tests (push) Has been cancelled
pytest / Integration (integration_tests_a) (push) Has been cancelled
pytest / Integration (integration_tests_b) (push) Has been cancelled
pytest / Integration (integration_tests_c) (push) Has been cancelled
pytest / Integration (integration_tests_d) (push) Has been cancelled
pytest / Integration (integration_tests_e) (push) Has been cancelled
pytest / Integration (integration_tests_f) (push) Has been cancelled
pytest / Integration (integration_tests_g) (push) Has been cancelled
pytest / Integration (integration_tests_h) (push) Has been cancelled
pytest / Integration (integration_tests_i) (push) Has been cancelled
pytest / Integration (integration_tests_j) (push) Has been cancelled
pytest / Distributed (distributed_a) (push) Has been cancelled
pytest / Distributed (distributed_b) (push) Has been cancelled
pytest / Distributed (distributed_c) (push) Has been cancelled
pytest / Distributed (distributed_d) (push) Has been cancelled
pytest / Distributed (distributed_e) (push) Has been cancelled
pytest / Distributed (distributed_f) (push) Has been cancelled
pytest / Minimal Install (push) Has been cancelled
pytest / Event File (push) Has been cancelled
pytest (slow) / py-slow (push) Has been cancelled
Publish JSON Schema / publish-schema (push) Has been cancelled
197 行
7.2 KiB
Python
197 行
7.2 KiB
Python
from typing import Any
|
|
|
|
from ludwig.api_annotations import DeveloperAPI
|
|
from ludwig.constants import (
|
|
DECODER,
|
|
ENCODER,
|
|
IMAGE,
|
|
INPUT_FEATURES,
|
|
MODEL_ECD,
|
|
MODEL_LLM,
|
|
MODEL_TYPE,
|
|
PREPROCESSING,
|
|
SEQUENCE,
|
|
TEXT,
|
|
TIMESERIES,
|
|
TYPE,
|
|
)
|
|
from ludwig.features.feature_registries import get_input_type_registry
|
|
from ludwig.schema.model_config import ModelConfig
|
|
from ludwig.types import FeatureConfigDict, FeatureTypeDefaultsDict, PreprocessingConfigDict
|
|
|
|
|
|
@DeveloperAPI
|
|
def get_feature_type_parameter_values_from_section(
|
|
config: ModelConfig, features_section: str, feature_type: str, parameter_name: str
|
|
) -> set:
|
|
"""Returns the set of all parameter values used for the given features_section, feature_type, and
|
|
parameter_name."""
|
|
parameter_values = set()
|
|
for feature in config[features_section]:
|
|
if feature[TYPE] == feature_type:
|
|
if parameter_name in feature:
|
|
parameter_values.add(feature[parameter_name])
|
|
elif parameter_name in feature[ENCODER]:
|
|
parameter_values.add(feature[ENCODER][parameter_name])
|
|
elif parameter_name in feature[DECODER]:
|
|
parameter_values.add(feature[DECODER][parameter_name])
|
|
return parameter_values
|
|
|
|
|
|
@DeveloperAPI
|
|
def get_defaults_section_for_feature_type(
|
|
feature_type: str,
|
|
config_defaults: FeatureTypeDefaultsDict,
|
|
config_defaults_section: str,
|
|
) -> FeatureConfigDict:
|
|
"""Returns a dictionary of all default parameter values specified in the global defaults section for the
|
|
config_defaults_section of the feature_type."""
|
|
|
|
if feature_type not in config_defaults:
|
|
return {}
|
|
|
|
if config_defaults_section not in config_defaults[feature_type]:
|
|
return {}
|
|
|
|
return config_defaults[feature_type][config_defaults_section]
|
|
|
|
|
|
def _to_dict(obj) -> dict:
|
|
"""Convert a config object or dict to a plain dict."""
|
|
if isinstance(obj, dict):
|
|
return obj
|
|
return obj.to_dict()
|
|
|
|
|
|
def get_preprocessing_params(config_obj: ModelConfig) -> PreprocessingConfigDict:
|
|
"""Returns a new dictionary that merges preprocessing section of config with type-specific preprocessing
|
|
parameters from config defaults."""
|
|
preprocessing_params = {}
|
|
preprocessing_params.update(_to_dict(config_obj.preprocessing))
|
|
for feat_type in get_input_type_registry():
|
|
if hasattr(config_obj.defaults, feat_type):
|
|
feat_defaults = getattr(config_obj.defaults, feat_type)
|
|
preprocessing = (
|
|
feat_defaults.preprocessing
|
|
if not isinstance(feat_defaults, dict)
|
|
else feat_defaults.get("preprocessing", {})
|
|
)
|
|
preprocessing_params[feat_type] = _to_dict(preprocessing)
|
|
return preprocessing_params
|
|
|
|
|
|
@DeveloperAPI
|
|
def merge_config_preprocessing_with_feature_specific_defaults(
|
|
config_preprocessing: PreprocessingConfigDict, config_defaults: FeatureTypeDefaultsDict
|
|
) -> PreprocessingConfigDict:
|
|
"""Returns a new dictionary that merges preprocessing section of config with type-specific preprocessing
|
|
parameters from config defaults."""
|
|
preprocessing_params = {}
|
|
preprocessing_params.update(config_preprocessing)
|
|
for feature_type in config_defaults:
|
|
preprocessing_params[feature_type] = config_defaults[feature_type].get(PREPROCESSING, {})
|
|
return preprocessing_params
|
|
|
|
|
|
def has_trainable_encoder(config: ModelConfig) -> bool:
|
|
for feature in config.input_features.to_list():
|
|
encoder = feature.get("encoder", {})
|
|
if encoder.get("trainable", False):
|
|
# TODO(travis): we assume here that False is always the default, which may not be true. We should dervice
|
|
# this from the schema.
|
|
return True
|
|
|
|
return False
|
|
|
|
|
|
def has_unstructured_input_feature(config: ModelConfig) -> bool:
|
|
for feature in config.input_features.to_list():
|
|
if feature.get("type", None) in {TEXT, IMAGE, SEQUENCE, TIMESERIES}:
|
|
return True
|
|
return False
|
|
|
|
|
|
def has_pretrained_encoder(config: ModelConfig) -> bool:
|
|
for feature in config.input_features:
|
|
if feature.encoder.is_pretrained():
|
|
return True
|
|
return False
|
|
|
|
|
|
def config_uses_llm(config: dict[str, Any] | ModelConfig) -> bool:
|
|
"""Determine if a config uses an LLM.
|
|
|
|
Args:
|
|
config: Ludwig config object or dictionary
|
|
|
|
Returns:
|
|
True if the model type is LLM or if the model uses and LLM encoder, otherwise False.
|
|
"""
|
|
uses_llm = False
|
|
|
|
# For a valid config, model_type LLM is automatically True
|
|
# ECD models need to be checked for at least one LLM text encoder
|
|
if isinstance(config, ModelConfig):
|
|
if config.model_type == MODEL_LLM:
|
|
uses_llm = True
|
|
else:
|
|
for feature in config.input_features:
|
|
if feature.encoder and feature.encoder.type == MODEL_LLM:
|
|
uses_llm = True
|
|
break
|
|
elif isinstance(config, dict) and config:
|
|
if config.get(MODEL_TYPE, MODEL_ECD) == MODEL_LLM:
|
|
uses_llm = True
|
|
elif INPUT_FEATURES in config:
|
|
for feature in config.get(INPUT_FEATURES, []):
|
|
if feature.get(ENCODER, {}).get(TYPE) == MODEL_LLM:
|
|
uses_llm = True
|
|
break
|
|
else:
|
|
raise ValueError(
|
|
f"Invalid config cannot be checked for LLM usage because it has no input features.Config: {config}"
|
|
)
|
|
else:
|
|
raise ValueError(f"Invalid config cannot be checked for LLM usage. Config: {config}")
|
|
|
|
return uses_llm
|
|
|
|
|
|
def get_quantization(config: dict[str, Any] | ModelConfig) -> list[int | None]:
|
|
"""Get the quantization specified in a config at any level.
|
|
|
|
Args:
|
|
config: Ludwig config object or dictionary
|
|
|
|
Returns:
|
|
For LLM models, the value of quantization.bits or None if it is not specified.
|
|
For ECD models, the list of values of quantization.bits for each encoder. If the encoder does not
|
|
support quantization or no quantization config is specified, the list entry is None.
|
|
"""
|
|
if isinstance(config, ModelConfig):
|
|
if config.model_type == MODEL_LLM:
|
|
return [config.quantization.bits] if config.quantization else [None]
|
|
else:
|
|
quantization_bits = []
|
|
for feature in config.input_features:
|
|
try:
|
|
quantization = feature.encoder.quantization.bits
|
|
except AttributeError:
|
|
quantization = None
|
|
quantization_bits.append(quantization)
|
|
return quantization_bits
|
|
elif isinstance(config, dict) and config:
|
|
if config.get(MODEL_TYPE, MODEL_ECD) == MODEL_LLM:
|
|
return [config.get("quantization", {}).get("bits")]
|
|
elif INPUT_FEATURES in config:
|
|
quantization_bits = []
|
|
for feature in config.get(INPUT_FEATURES, []):
|
|
quantization_bits.append(feature.get(ENCODER, {}).get("quantization", {}).get("bits"))
|
|
return quantization_bits
|
|
else:
|
|
raise ValueError(
|
|
f"Invalid config cannot be checked for quantization because it has no input features.Config: {config}"
|
|
)
|
|
else:
|
|
raise ValueError(f"Invalid config cannot be checked for quantization. Config: {config}")
|