Skip to content
Merged
79 changes: 57 additions & 22 deletions lm_scorer/models/abc/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,38 +10,73 @@ class LMScorer(ABC):
def __init__(self, model_name: str, **kwargs: Any) -> None:
self._build(model_name, kwargs)

@overload
def sentence_score(
self, text: str, log: bool = False, reduce: str = "prod"
) -> float:
log_probs, _, _ = self._tokens_log_prob(text)
tlen = log_probs.shape[0]
...

if reduce == "prod":
score = log_probs.sum()
elif reduce == "mean":
score = log_probs.logsumexp(0) - math.log(tlen)
elif reduce == "gmean":
score = log_probs.mean(0)
elif reduce == "hmean":
score = log_probs.neg().logsumexp(0).neg() + math.log(tlen)
else:
raise ValueError("Unrecognized scoring strategy: %s" % reduce)
@overload
def sentence_score(
self, text: List[str], log: bool = False, reduce: str = "prod"
) -> List[float]:
...

def sentence_score(
self, text: Union[str, List[str]], log: bool = False, reduce: str = "prod",
) -> Union[float, List[float]]:
sentences = [text] if isinstance(text, str) else text
outputs = self._tokens_log_prob(sentences)

scores = []
for output in outputs:
log_probs = output[0]
tlen = log_probs.shape[0]

if reduce == "prod":
score = log_probs.sum()
elif reduce == "mean":
score = log_probs.logsumexp(0) - math.log(tlen)
elif reduce == "gmean":
score = log_probs.mean(0)
elif reduce == "hmean":
score = log_probs.neg().logsumexp(0).neg() + math.log(tlen)
else:
raise ValueError("Unrecognized scoring strategy: %s" % reduce)

if not log:
score = score.exp()
if not log:
score = score.exp()

return score.item()
scores.append(score.item())

return scores[0] if isinstance(text, str) else scores

@overload
def tokens_score(
self, text: str, log: bool = False
) -> Tuple[List[float], List[int], List[str]]:
Comment thread
dldk-gael marked this conversation as resolved.
log_probs, ids, tokens = self._tokens_log_prob(text)
...

@overload
def tokens_score(
self, text: List[str], log: bool = False
) -> List[Tuple[List[float], List[int], List[str]]]:
...

scores = log_probs # type: torch.Tensor # type: ignore
if not log:
scores = scores.exp()
def tokens_score(
self, text: Union[str, List[str]], log: bool = False
) -> Union[
Tuple[List[float], List[int], List[str]],
List[Tuple[List[float], List[int], List[str]]],
]:
sentences = [text] if isinstance(text, str) else text
outputs = []
for scores, ids, tokens in self._tokens_log_prob(sentences):
Comment thread
dldk-gael marked this conversation as resolved.
Outdated
if not log:
scores = scores.exp() # type: torch.Tensor # type: ignore
Comment thread
dldk-gael marked this conversation as resolved.
Outdated
outputs.append((scores.tolist(), ids.tolist(), tokens))

return scores.tolist(), ids.tolist(), tokens
return outputs[0] if isinstance(text, str) else outputs

@classmethod
def supported_model_names(cls) -> Iterable[str]:
Expand All @@ -53,8 +88,8 @@ def _build(self, model_name: str, options: Dict[str, Any]) -> None:

@abstractmethod
def _tokens_log_prob(
self, text: str
) -> Tuple[torch.FloatTensor, torch.LongTensor, List[str]]:
self, sentences: List[str]
Comment thread
dldk-gael marked this conversation as resolved.
Outdated
) -> List[Tuple[torch.FloatTensor, torch.LongTensor, List[str]]]:
... # pragma: no cover

@classmethod
Expand Down
94 changes: 69 additions & 25 deletions lm_scorer/models/gpt2.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from typing import * # pylint: disable=wildcard-import,unused-wildcard-import


import torch
from transformers import GPT2Tokenizer, GPT2LMHeadModel

Expand All @@ -18,46 +19,89 @@ def _build(self, model_name: str, options: Dict[str, Any]) -> None:
if "device" in options:
self.model.to(options["device"])

# @overrides
def _tokens_log_prob(
self, text: str
) -> Tuple[torch.FloatTensor, torch.LongTensor, List[str]]:
device = self.model.device
input_text = "%s%s%s" % (
self.tokenizer.bos_token,
text,
self.tokenizer.eos_token,
self.batch_size = options["batch_size"] if "batch_size" in options else 1

def add_special_tokens_and_encode(self, text):
return self.tokenizer.encode(
self.tokenizer.bos_token + text + self.tokenizer.eos_token
)
# len(tokens) = seq_len + 2
tokens = self.tokenizer.tokenize(input_text)
# ids.shape = [1, seq_len + 2]
ids = torch.tensor( # pylint: disable=not-callable
[self.tokenizer.convert_tokens_to_ids(tokens)],
device=device,
dtype=torch.long,

def pad(self, sequences: List[torch.Tensor]):
max_seq_len = max([s.size(0) for s in sequences])
out_tensor = (
sequences[0]
.data.new_zeros(len(sequences), max_seq_len)
.fill_(self.tokenizer.eos_token_id)
)
mask = torch.zeros((len(sequences), max_seq_len), device=sequences[0].device)
for i, tensor in enumerate(sequences):
length = tensor.size(0)
out_tensor[i, :length] = tensor
mask[i, :length] = 1

return out_tensor, mask

def _tokens_log_prob_single_batch(
self, sentences: List[str]
) -> List[Tuple[torch.FloatTensor, torch.LongTensor, List[str]]]:

device = self.model.device

tokens = [
self.add_special_tokens_and_encode(sentence) for sentence in sentences
]
ids, mask = self.pad(
list(
map(
lambda x: torch.tensor( # pylint: disable=not-callable
x, device=device, dtype=torch.long
),
tokens,
)
)
)

with torch.no_grad():
outputs = self.model(ids)

# pred_scores.shape = [1, seq_len + 2, vocab_size]
# pred_scores.shape = [nb_sentences, max_seq_len, vocab_size]
pred_scores = outputs[0]

# len(tokens) = seq_len + 1
tokens = tokens[1:]
# ids.shape = [1, seq_len + 1, vocab_size]
# Align input and target
ids = ids[:, 1:]
# pred_scores.shape = [1, seq_len + 1, vocab_size]
pred_scores = pred_scores[:, :-1, :]

# ids_scores.shape = [1, seq_len + 1]
# Retrieve the token scores corresponding to the target id
# ids_scores.shape = [nb_sentences, max_seq_len]
ids_scores = pred_scores.gather(2, ids.unsqueeze(2)).squeeze(2)
# log_prob.shape = [1, seq_len + 1]

# Zero the values of the padding inputs
ids_scores *= mask[:, 1:]

# log_prob.shape = [nb_sentences, max_seq_len]
log_probs = ids_scores - pred_scores.logsumexp(2)

return log_probs[0], ids[0], tokens # type: ignore
return [
(log_probs[i, : len(tokens[i])], ids[i, : len(tokens[i])], tokens[i][1:])
for i in range(len(sentences))
] # type: ignore

# @overrides
def _tokens_log_prob(
self, sentences: List[str]
Comment thread
dldk-gael marked this conversation as resolved.
Outdated
) -> List[Tuple[torch.FloatTensor, torch.LongTensor, List[str]]]:

output = []
for i in range(len(sentences) // self.batch_size):
Comment thread
dldk-gael marked this conversation as resolved.
Outdated
output += self._tokens_log_prob_single_batch(
sentences[i * self.batch_size : (i + 1) * self.batch_size]
)
if len(sentences) % self.batch_size != 0:
output += self._tokens_log_prob_single_batch(
sentences[-(len(sentences) % self.batch_size) :]
)
return output

# @overrides_tokens_log_prob_single_batch
@classmethod
def _supported_model_names(cls) -> Iterable[str]:
return GPT2LMHeadModel.pretrained_model_archive_map.keys()