From 53f66172865193869691d3a1f0c4dc89e7ae2d56 Mon Sep 17 00:00:00 2001 From: Richard Rogers Date: Thu, 26 Oct 2023 09:42:21 +0000 Subject: [PATCH] experimental multilingual idea --- langkit/all_metrics.py | 23 ++++++++++---------- langkit/count_regexes.py | 11 ++++++---- langkit/injections.py | 32 ++++++++++++++++------------ langkit/input_output.py | 38 +++++++++++++++++++-------------- langkit/light_metrics.py | 11 +++++----- langkit/llm_metrics.py | 19 +++++++++-------- langkit/nlp_scores.py | 14 ++++++++----- langkit/regexes.py | 8 ++++--- langkit/sentiment.py | 14 ++++++++++--- langkit/textstat.py | 45 ++++++++++++++++++++++------------------ langkit/themes.py | 19 ++++++++++------- langkit/topics.py | 7 ++++++- langkit/toxicity.py | 14 ++++++++++--- 13 files changed, 155 insertions(+), 100 deletions(-) diff --git a/langkit/all_metrics.py b/langkit/all_metrics.py index cf3198ab..1447c388 100644 --- a/langkit/all_metrics.py +++ b/langkit/all_metrics.py @@ -1,4 +1,4 @@ -from typing import Optional +from typing import List, Optional from whylogs.experimental.core.udf_schema import udf_schema from whylogs.core.schema import DeclarativeSchema @@ -13,14 +13,15 @@ from langkit import input_output -def init(config: Optional[LangKitConfig] = None) -> DeclarativeSchema: - injections.init(config=config) - topics.init(config=config) - regexes.init(config=config) - sentiment.init(config=config) - textstat.init(config=config) - themes.init(config=config) - toxicity.init(config=config) - input_output.init(config=config) - text_schema = udf_schema() +def init(languages: List[str] = ["en"], config: Optional[LangKitConfig] = None) -> DeclarativeSchema: + for language in langauges: + injections.init(language, config=config) + topics.init(language, config=config) + regexes.init(language, config=config) + sentiment.init(language, config=config) + textstat.init(language, config=config) + themes.init(language, config=config) + toxicity.init(language, config=config) + input_output.init(language, config=config) + text_schema = udf_schema(chained_schemas=languages) return text_schema diff --git a/langkit/count_regexes.py b/langkit/count_regexes.py index c9b1ef53..f06a1b8b 100644 --- a/langkit/count_regexes.py +++ b/langkit/count_regexes.py @@ -44,23 +44,26 @@ def _unregister(): _registered = set() -def _register_udfs(): +def _register_udfs(language: str): global _registered _unregister() regex_groups = pattern_loader.get_regex_groups() if regex_groups is not None: for column in [prompt_column, response_column]: for group in regex_groups: - udf_name = f"{column}.{group['name']}_count" + udf_name = f"{language}.{column}.{group['name']}_count" register_dataset_udf( [column], udf_name=udf_name, + schema_name=language )(wrapper(group, column)) _registered.add(udf_name) def init( - pattern_file_path: Optional[str] = None, config: Optional[LangKitConfig] = None + language: str = "en", + pattern_file_path: Optional[str] = None, + config: Optional[LangKitConfig] = None ): config = deepcopy(config or lang_config) if pattern_file_path: @@ -70,7 +73,7 @@ def init( pattern_loader = PatternLoader(config) pattern_loader.update_patterns() - _register_udfs() + _register_udfs(language) init() diff --git a/langkit/injections.py b/langkit/injections.py index 26f55111..69ffe837 100644 --- a/langkit/injections.py +++ b/langkit/injections.py @@ -23,7 +23,21 @@ def download_embeddings(url): return array +def injection(prompt: Union[Dict[str, List], pd.DataFrame]) -> Union[List, pd.Series]: + global _transformer_model + global _index_embeddings + if _transformer_model is None: + raise ValueError("Injections - transformer model not initialized") + embeddings = _transformer_model.encode(prompt[_prompt]) + faiss.normalize_L2(embeddings) + if _index_embeddings is None: + raise ValueError("Injections - index embeddings not initialized") + dists, _ = _index_embeddings.search(x=embeddings, k=1) + return dists.flatten().tolist() + + def init( + language: str = "en", transformer_name: Optional[str] = None, version: Optional[str] = None, config: Optional[LangKitConfig] = None, @@ -73,19 +87,11 @@ def init( f"Injections - unable to deserialize index to {embeddings_path}. Error: {deserialization_error}" ) - -@register_dataset_udf([_prompt], f"{_prompt}.injection") -def injection(prompt: Union[Dict[str, List], pd.DataFrame]) -> Union[List, pd.Series]: - global _transformer_model - global _index_embeddings - if _transformer_model is None: - raise ValueError("Injections - transformer model not initialized") - embeddings = _transformer_model.encode(prompt[_prompt]) - faiss.normalize_L2(embeddings) - if _index_embeddings is None: - raise ValueError("Injections - index embeddings not initialized") - dists, _ = _index_embeddings.search(x=embeddings, k=1) - return dists.flatten().tolist() + register_dataset_udf( + [_prompt], + udf_name=f"{language}{_prompt}.injection", + schema_name=language + )(injection) init() diff --git a/langkit/input_output.py b/langkit/input_output.py index 766798b7..c8e36587 100644 --- a/langkit/input_output.py +++ b/langkit/input_output.py @@ -16,22 +16,6 @@ diagnostic_logger = getLogger(__name__) -def init( - transformer_name: Optional[str] = None, - custom_encoder: Optional[Callable] = None, - config: Optional[LangKitConfig] = None, -): - config = config or deepcopy(lang_config) - global _transformer_model - if transformer_name is None and custom_encoder is None: - transformer_name = config.transformer_name - _transformer_model = Encoder(transformer_name, custom_encoder) - - -init() - - -@register_dataset_udf([_prompt, _response], f"{_response}.relevance_to_{_prompt}") def prompt_response_similarity(text): global _transformer_model @@ -53,3 +37,25 @@ def prompt_response_similarity(text): ) series_result.append(None) return series_result + + +def init( + language: str = "en" + transformer_name: Optional[str] = None, + custom_encoder: Optional[Callable] = None, + config: Optional[LangKitConfig] = None, +): + config = config or deepcopy(lang_config) + global _transformer_model + if transformer_name is None and custom_encoder is None: + transformer_name = config.transformer_name + _transformer_model = Encoder(transformer_name, custom_encoder) + register_dataset_udf( + [_prompt, _response], + f"{language}.{_response}.relevance_to_{_prompt}", + schema_name=language + )(prompt_response_similarity) + + +init() + diff --git a/langkit/light_metrics.py b/langkit/light_metrics.py index c0779e72..85283e81 100644 --- a/langkit/light_metrics.py +++ b/langkit/light_metrics.py @@ -1,4 +1,4 @@ -from typing import Optional +from typing import List, Optional from whylogs.experimental.core.udf_schema import udf_schema from whylogs.core.schema import DeclarativeSchema @@ -7,9 +7,10 @@ from langkit import textstat -def init(config: Optional[LangKitConfig] = None) -> DeclarativeSchema: - regexes.init(config=config) - textstat.init(config=config) +def init(languages: List[str] = ["en"], config: Optional[LangKitConfig] = None) -> DeclarativeSchema: + for language in languages: + regexes.init(language, config=config) + textstat.init(language, config=config) - text_schema = udf_schema() + text_schema = udf_schema(chained_schemas=languages) return text_schema diff --git a/langkit/llm_metrics.py b/langkit/llm_metrics.py index ea44d5e1..f4e331f3 100644 --- a/langkit/llm_metrics.py +++ b/langkit/llm_metrics.py @@ -1,6 +1,6 @@ from . import LangKitConfig from logging import getLogger -from typing import Optional +from typing import List, Optional from whylogs.experimental.core.udf_schema import udf_schema from whylogs.core.schema import DeclarativeSchema @@ -19,13 +19,14 @@ ) -def init(config: Optional[LangKitConfig] = None) -> DeclarativeSchema: - regexes.init(config=config) - sentiment.init(config=config) - textstat.init(config=config) - themes.init(config=config) - toxicity.init(config=config) - input_output.init(config=config) +def init(languages: List[str] = ["en"], config: Optional[LangKitConfig] = None) -> DeclarativeSchema: + for language in languages: + regexes.init(language, config=config) + sentiment.init(language, config=config) + textstat.init(language, config=config) + themes.init(language, config=config) + toxicity.init(language, config=config) + input_output.init(language, config=config) - text_schema = udf_schema() + text_schema = udf_schema(chained_schemas = languages) return text_schema diff --git a/langkit/nlp_scores.py b/langkit/nlp_scores.py index 3bd6fa77..268a3b60 100644 --- a/langkit/nlp_scores.py +++ b/langkit/nlp_scores.py @@ -19,7 +19,7 @@ _meteor_registered = False -def _register_score_udfs(): +def _register_score_udfs(language: str): global _bleu_registered, _rouge_registered, _meteor_registered if _corpus: @@ -30,7 +30,8 @@ def _register_score_udfs(): @register_dataset_udf( [response_column], - udf_name=f"{response_column}.bleu_score", + udf_name=f"{language}.{response_column}.bleu_score", + schema_name=language ) def bleu_score(text): result = [] @@ -48,7 +49,8 @@ def bleu_score(text): @register_dataset_udf( [response_column], - udf_name=f"{response_column}.rouge_score", + udf_name=f"{language}.{response_column}.rouge_score", + schema_name=language ) def rouge_score(text): result = [] @@ -68,7 +70,8 @@ def rouge_score(text): @register_dataset_udf( [response_column], - udf_name=f"{response_column}.meteor_score", + udf_name=f"{language}.{response_column}.meteor_score", + schema_name=language ) def meteor_score(text): result = [] @@ -87,6 +90,7 @@ def meteor_score(text): def init( + language: str = "en", corpus: Optional[str] = None, scores: Set[str] = set(), rouge_type: str = "", @@ -100,7 +104,7 @@ def init( _scores = list(scores or config.nlp_scores) _rouge_type = rouge_type or config.rouge_type - _register_score_udfs() + _register_score_udfs(language) init() diff --git a/langkit/regexes.py b/langkit/regexes.py index a8cf10ee..d8deaa38 100644 --- a/langkit/regexes.py +++ b/langkit/regexes.py @@ -38,7 +38,7 @@ def wrappee(text): _registered = False -def _register_udfs(): +def _register_udfs(language: str): global _registered if _registered: return @@ -48,12 +48,14 @@ def _register_udfs(): for column in [prompt_column, response_column]: register_dataset_udf( [column], - udf_name=f"{column}.has_patterns", + udf_name=f"{language}.{column}.has_patterns", + schema_name=language, metrics=[MetricSpec(FrequentItemsMetric)], )(_wrapper(column)) def init( + language: str = "en", pattern_file_path: Optional[str] = None, config: Optional[LangKitConfig] = None ): config = deepcopy(config or lang_config) @@ -64,7 +66,7 @@ def init( pattern_loader = PatternLoader(config) pattern_loader.update_patterns() - _register_udfs() + _register_udfs(language) init() diff --git a/langkit/sentiment.py b/langkit/sentiment.py index cadf7ffe..e91d25af 100644 --- a/langkit/sentiment.py +++ b/langkit/sentiment.py @@ -19,17 +19,15 @@ def sentiment_nltk(text: str) -> float: return _sentiment_analyzer.polarity_scores(text)["compound"] -@register_dataset_udf([_prompt], udf_name=f"{_prompt}.sentiment_nltk") def prompt_sentiment(text): return [sentiment_nltk(t) for t in text[_prompt]] -@register_dataset_udf([_response], udf_name=f"{_response}.sentiment_nltk") def response_sentiment(text): return [sentiment_nltk(t) for t in text[_response]] -def init(lexicon: Optional[str] = None, config: Optional[LangKitConfig] = None): +def init(language: str = "en", lexicon: Optional[str] = None, config: Optional[LangKitConfig] = None): import nltk from nltk.sentiment import SentimentIntensityAnalyzer @@ -41,6 +39,16 @@ def init(lexicon: Optional[str] = None, config: Optional[LangKitConfig] = None): _nltk_downloaded = True _sentiment_analyzer = SentimentIntensityAnalyzer() + register_dataset_udf( + [_prompt], + udf_name=f"{language}.{_prompt}.sentiment_nltk", + schema_name=language + )(prompt_sentiment) + register_dataset_udf( + [_response], + udf_name=f"{language}.{_response}.sentiment_nltk", + schema_name=language + )(response_sentiment) init() diff --git a/langkit/textstat.py b/langkit/textstat.py index 9ca0d678..26904d29 100644 --- a/langkit/textstat.py +++ b/langkit/textstat.py @@ -62,31 +62,36 @@ def wrappee(text: Union[pd.DataFrame, Dict[str, List]]) -> Union[pd.Series, List return wrappee -def init(config: Optional[LangKitConfig] = None): - pass - - -init() - - def _unpack(t: Union[Tuple[str, str], Tuple[str, str, str]]) -> Tuple[str, str, str]: return t if len(t) == 3 else (t[0], t[1], t[0]) # type: ignore -_registered = False +_registered: Dict[str, bool] = dict() -if not _registered: - _registered = True - for t in _udfs_to_register: - stat_name, schema_name, udf = _unpack(t) +def init(language: str = "en", config: Optional[LangKitConfig] = None): + global _registered + if not _registered.get(language, False): + _registered[language] = True + for t in _udfs_to_register: + stat_name, schema_name, udf = _unpack(t) + language = schema_name or language # TODO: double-check this + for column in [prompt_column, response_column]: + register_dataset_udf( + [column], + udf_name=f"{language}.{column}.{udf}", + schema_name=language + )(wrapper(stat_name, column)) for column in [prompt_column, response_column]: register_dataset_udf( - [column], udf_name=f"{column}.{udf}", schema_name=schema_name - )(wrapper(stat_name, column)) - for column in [prompt_column, response_column]: - register_dataset_udf([column], udf_name=f"{column}.aggregate_reading_level")( - aggregate_wrapper(column) - ) - - diagnostic_logger.info("Initialized textstat metrics.") + [column], + udf_name=f"{language}.{column}.aggregate_reading_level", + schema_name=language + )( + aggregate_wrapper(column) + ) + + diagnostic_logger.info("Initialized textstat metrics.") + + +init() diff --git a/langkit/themes.py b/langkit/themes.py index 484b24db..6b6ead0c 100644 --- a/langkit/themes.py +++ b/langkit/themes.py @@ -52,10 +52,10 @@ def _map_embeddings(): ] -_registered = set() +_registered: Dict[str, Set] = set() -def _register_theme_udfs(): +def _register_theme_udfs(language: str): global _registered _map_embeddings() @@ -65,10 +65,14 @@ def _register_theme_udfs(): continue if group == "refusal" and column == _prompt: continue - udf_name = f"{column}.{group}_similarity" - if udf_name not in _registered: - _registered.add(udf_name) - register_dataset_udf([column], udf_name=udf_name)( + udf_name = f"{language}.{column}.{group}_similarity" + if udf_name not in _registered.get(language, set()): + _registered[language].add(udf_name) # TODO: use defaultdict + register_dataset_udf( + [column], + udf_name=udf_name, + schema_name=language + )( create_similarity_function(group, column) ) @@ -90,6 +94,7 @@ def load_themes(json_path: str, encoding="utf-8"): def init( + language: str = "en", transformer_name: Optional[str] = None, custom_encoder: Optional[Callable] = None, theme_file_path: Optional[str] = None, @@ -111,7 +116,7 @@ def init( _theme_groups = load_themes(config.theme_file_path) else: _theme_groups = load_themes(theme_file_path) - _register_theme_udfs() + _register_theme_udfs(language) def get_subject_similarity(text: str, comparison_embedding: Tensor) -> float: diff --git a/langkit/topics.py b/langkit/topics.py index 598d8992..f382b613 100644 --- a/langkit/topics.py +++ b/langkit/topics.py @@ -22,6 +22,7 @@ def _wrapper(column: str) -> Callable: def init( + language: str = "en", topics: Optional[List[str]] = None, model_path: Optional[str] = None, topic_classifier: Optional[str] = None, @@ -34,7 +35,11 @@ def init( model_path = model_path or config.topic_model_path _classifier = pipeline(topic_classifier, model=model_path) for column in [prompt_column, response_column]: - register_dataset_udf([column], udf_name=f"{column}.closest_topic")( + register_dataset_udf( + [column], + udf_name=f"{language}.{column}.closest_topic", + schema_name=language + )( _wrapper(column) ) diff --git a/langkit/toxicity.py b/langkit/toxicity.py index 9dd31623..b4c9861e 100644 --- a/langkit/toxicity.py +++ b/langkit/toxicity.py @@ -23,17 +23,15 @@ def toxicity(text: str) -> float: ) -@register_dataset_udf([_prompt], f"{_prompt}.toxicity") def prompt_toxicity(text): return [toxicity(t) for t in text[_prompt]] -@register_dataset_udf([_response], f"{_response}.toxicity") def response_toxicity(text): return [toxicity(t) for t in text[_response]] -def init(model_path: Optional[str] = None, config: Optional[LangKitConfig] = None): +def init(language: str = "en", model_path: Optional[str] = None, config: Optional[LangKitConfig] = None): from transformers import ( AutoModelForSequenceClassification, AutoTokenizer, @@ -48,6 +46,16 @@ def init(model_path: Optional[str] = None, config: Optional[LangKitConfig] = Non _toxicity_pipeline = TextClassificationPipeline( model=model, tokenizer=_toxicity_tokenizer ) + register_dataset_udf( + [_prompt], + f"{language}.{_prompt}.toxicity", + schema_name=language + )(prompt_toxicity) + register_dataset_udf( + [_response], + f"{language}.{_response}.toxicity" + schema_name=language + )(response_toxicity) init()