AVG / language_model.py
anonymous5378's picture
Fix model name
9bf6109 verified
Raw
History Blame
8 kB
import os
import re
import logging
import pandas as pd
import torch
import transformers
from transformers import AutoTokenizer, AutoConfig, T5ForConditionalGeneration, BartForConditionalGeneration
TASK_PREFIX = {
"ve": "extract values",
"ag": "generate attributes",
"avg": "generate attributes and values",
"av": "attribute value extraction"
}
ADDITIONAL_SP_TOKENS = {'hl': '<hl>'}
def load_language_model(model_id,
):
""" Load language model from hugging face hub. """
tokenizer = AutoTokenizer.from_pretrained(model_id)
config = AutoConfig.from_pretrained(model_id)
# model class
if config.model_type == 't5':
model_class = T5ForConditionalGeneration.from_pretrained
elif config.model_type == 'bart':
model_class = BartForConditionalGeneration.from_pretrained
else:
raise ValueError(f"Unsupported model type: {config.model_type}")
param = {'config': config}
model = model_class(model_id, **param)
return tokenizer, model, config
class TransformersAVG:
""" Transformers Language Model for Attribute Value Generation. """
def __init__(self,
model: str = None,
max_input_length: int = 512,
max_target_length: int = 256,
model_ve: str = None,
max_target_length_ve: int = 34,
is_ag: bool = None,
is_avg: bool = None,
is_ve: bool = None,
is_av: bool = None
) -> None:
self.is_ag = 'ag' in model.split('-') if is_ag is None else is_ag
self.is_ve = 've' in model.split('-') if is_ve is None else is_ve
self.is_av = 'mlt' in model.split('-') if is_av is None else is_av
self.is_avg = 'end2end' in model.split('-') if is_avg is None else is_avg
self.model_name = model
self.max_input_length = max_input_length
self.max_target_length = max_target_length
self.model_name_ve = model_ve
self.max_target_length_ve = max_target_length_ve
# load model
self.tokenizer, self.model, config = load_language_model(
self.model_name,
)
# Setup GPU device
self.device = 'cuda' if torch.cuda.device_count() > 0 else 'cpu'
self.model.to(self.device)
def generate_av_end2end(self,
context: str,
num_beams: int = 4,
splitting_symbol: str = '|'
):
""" Generate attribute value pairs in an end2end fashion. """
logging.info(f"running model for 'attribute value pair generation'.")
model_input = self.tokenizer(context, max_length=self.max_input_length, truncation=True,
padding="max_length")
model_input = {k:torch.unsqueeze(torch.tensor(v),dim=0) for k,v in model_input.items()}
outputs = self.model.generate(**model_input, num_beams=num_beams, do_sample=True, max_length=self.max_target_length)
outputs = self.tokenizer.batch_decode(outputs, skip_special_tokens=True)[0]
predictions = self.__format_av(outputs)
return predictions
def generate_av_pipeline(self,
context: str,
num_beams: int = 4
):
""" Generate attribute value pairs using the pipeline: first extract values then generate attributes """
logging.info(f"running model for value candidate extraction.")
# load ve model
self.ve_tokenizer, self.ve_model, ve_config = load_language_model(
self.model_name_ve,
)
# generate values
model_input = self.ve_tokenizer(context, max_length=self.max_input_length, truncation=True,
padding="max_length")
model_input = {k:torch.unsqueeze(torch.tensor(v),dim=0) for k,v in model_input.items()}
outputs = self.ve_model.generate(**model_input, num_beams=num_beams, do_sample=True, max_length=self.max_target_length_ve)
outputs = self.ve_tokenizer.batch_decode(outputs, skip_special_tokens=True)[0]
list_of_values = self.__format_v(outputs)
# generate attributes
list_of_attributes = []
logging.info(f"running model for attribute generation.")
list_of_contexts = self.__highlight_value(context, list_of_values)
model_input = self.tokenizer(list_of_contexts, max_length=self.max_input_length, truncation=True,
padding="max_length")
model_input = {k:torch.tensor(v) for k,v in model_input.items()}
outputs = self.model.generate(**model_input, num_beams=num_beams, do_sample=True, max_length=self.max_target_length)
outputs = self.tokenizer.batch_decode(outputs, skip_special_tokens=True)
list_of_attributes.append(outputs)
list_of_attributes = [pred for preds in list_of_attributes for pred in preds]
return [(att, val) for att, val in zip(list_of_attributes, list_of_values)]
def generate_av_mul(self,
context: str,
num_beams: int = 4):
""" Generate attribute-value pairs using a multi-task approach: same model for attribute genration and value extraction """
logging.info(f"running model for value extraction.")
model_input = self.tokenizer(f"extract value {context}", max_length=self.max_input_length, truncation=True,
padding="max_length")
model_input = {k:torch.unsqueeze(torch.tensor(v),dim=0) for k,v in model_input.items()}
outputs = self.model.generate(**model_input, num_beams=num_beams, do_sample=True, max_length=self.max_target_length)
outputs = self.tokenizer.batch_decode(outputs, skip_special_tokens=True)[0]
list_of_values = self.__format_v(outputs)
# generate attributes
list_of_attributes = []
logging.info(f"running model for attribute generation.")
list_of_contexts = self.__highlight_value(context, list_of_values, prefix=True)
model_input = self.tokenizer(list_of_contexts, max_length=self.max_input_length, truncation=True,
padding="max_length")
model_input = {k:torch.tensor(v) for k,v in model_input.items()}
outputs = self.model.generate(**model_input, num_beams=num_beams, do_sample=True, max_length=self.max_target_length)
outputs = self.tokenizer.batch_decode(outputs, skip_special_tokens=True)
list_of_attributes.append(outputs)
list_of_attributes = [pred for preds in list_of_attributes for pred in preds]
return [(att, val) for att, val in zip(list_of_attributes, list_of_values)]
def __highlight_value(self,
context: str,
list_of_values: list,
prefix: bool = False
):
list_of_contexts = []
for value in list_of_values:
my_context = context.replace(value, f'<hl> {value} <hl>')
if prefix:
my_context = f"generate attribute {my_context}"
list_of_contexts.append(my_context)
return list_of_contexts
def __format_v(self, predictions: str, splitting_symbol: str = '|') -> list:
my_list = []
for raw_string in predictions.split(splitting_symbol):
v = raw_string.strip()
my_list.append(v)
return my_list
def __format_av(self, predictions: str, splitting_symbol: str = '|') -> list:
my_list = []
for raw_string in predictions.split(splitting_symbol):
a = raw_string.replace('attribute: ', '')
a = a.split(',')[0].strip()
v = raw_string.replace('value: ', '')
v = v.split(',')[1].strip()
my_list.append((a,v))
return my_list