mirror of
https://github.com/wassname/discovering_latent_knowledge.git
synced 2026-09-12 12:13:04 +08:00
both pipelines seem to work
This commit is contained in:
+117
-43
@@ -87,12 +87,12 @@
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"\u001b[32m2023-10-24 10:45:58.826\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36msrc.models.load\u001b[0m:\u001b[36mverbose_change_param\u001b[0m:\u001b[36m18\u001b[0m - \u001b[1mchanging pad_token_id from 32000 to 0\u001b[0m\n",
|
||||
"2023-10-24T10:45:58.826830+0800 INFO changing pad_token_id from 32000 to 0\n",
|
||||
"\u001b[32m2023-10-24 10:45:58.827\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36msrc.models.load\u001b[0m:\u001b[36mverbose_change_param\u001b[0m:\u001b[36m18\u001b[0m - \u001b[1mchanging padding_side from right to left\u001b[0m\n",
|
||||
"2023-10-24T10:45:58.827930+0800 INFO changing padding_side from right to left\n",
|
||||
"\u001b[32m2023-10-24 10:45:58.828\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36msrc.models.load\u001b[0m:\u001b[36mverbose_change_param\u001b[0m:\u001b[36m18\u001b[0m - \u001b[1mchanging truncation_side from right to left\u001b[0m\n",
|
||||
"2023-10-24T10:45:58.828427+0800 INFO changing truncation_side from right to left\n"
|
||||
"\u001b[32m2023-10-24 10:49:20.938\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36msrc.models.load\u001b[0m:\u001b[36mverbose_change_param\u001b[0m:\u001b[36m18\u001b[0m - \u001b[1mchanging pad_token_id from 32000 to 0\u001b[0m\n",
|
||||
"2023-10-24T10:49:20.938690+0800 INFO changing pad_token_id from 32000 to 0\n",
|
||||
"\u001b[32m2023-10-24 10:49:20.940\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36msrc.models.load\u001b[0m:\u001b[36mverbose_change_param\u001b[0m:\u001b[36m18\u001b[0m - \u001b[1mchanging padding_side from right to left\u001b[0m\n",
|
||||
"2023-10-24T10:49:20.940322+0800 INFO changing padding_side from right to left\n",
|
||||
"\u001b[32m2023-10-24 10:49:20.941\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36msrc.models.load\u001b[0m:\u001b[36mverbose_change_param\u001b[0m:\u001b[36m18\u001b[0m - \u001b[1mchanging truncation_side from right to left\u001b[0m\n",
|
||||
"2023-10-24T10:49:20.941383+0800 INFO changing truncation_side from right to left\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
@@ -147,11 +147,11 @@
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"tokenize: 100%|██████████| 302/302 [00:00<00:00, 3370.75 examples/s]\n",
|
||||
"truncated: 100%|██████████| 302/302 [00:00<00:00, 3882.16 examples/s]\n",
|
||||
"prompt_truncated: 100%|██████████| 302/302 [00:00<00:00, 551.44 examples/s]\n",
|
||||
"choice_ids: 100%|██████████| 302/302 [00:00<00:00, 9920.74 examples/s]\n",
|
||||
"Filter: 100%|██████████| 302/302 [00:00<00:00, 3779.73 examples/s]"
|
||||
"tokenize: 100%|██████████| 302/302 [00:00<00:00, 3490.12 examples/s]\n",
|
||||
"truncated: 100%|██████████| 302/302 [00:00<00:00, 3889.90 examples/s]\n",
|
||||
"prompt_truncated: 100%|██████████| 302/302 [00:00<00:00, 604.08 examples/s]\n",
|
||||
"choice_ids: 100%|██████████| 302/302 [00:00<00:00, 9749.77 examples/s]\n",
|
||||
"Filter: 100%|██████████| 302/302 [00:00<00:00, 3987.01 examples/s]"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -236,7 +236,7 @@
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<src.repe.rep_readers.PCARepReader at 0x7f6ed001aa10>"
|
||||
"<src.repe.rep_readers.PCARepReader at 0x7fcd97b27250>"
|
||||
]
|
||||
},
|
||||
"execution_count": 8,
|
||||
@@ -264,7 +264,28 @@
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"{8: -7.1495447,\n",
|
||||
" 11: 16.796896,\n",
|
||||
" 14: 13.5725565,\n",
|
||||
" 17: -5.684828,\n",
|
||||
" 20: 26.571663,\n",
|
||||
" 23: 37.01544,\n",
|
||||
" 26: -67.76076,\n",
|
||||
" 29: -64.502205,\n",
|
||||
" 32: -52.72602,\n",
|
||||
" 35: -76.70331,\n",
|
||||
" 38: -80.74953}"
|
||||
]
|
||||
},
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"# read direction for each example, layer\n",
|
||||
"H_tests = rep_reading_pipeline(\n",
|
||||
@@ -285,14 +306,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
@@ -312,14 +326,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"execution_count": 11,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
@@ -329,7 +336,7 @@
|
||||
"\n",
|
||||
"# inputs = dataset_train[:2]\n",
|
||||
"# coeff=8.0\n",
|
||||
"# max_new_tokens=3\n",
|
||||
"# max_new_tokens=1\n",
|
||||
"# text_gen_kwargs = dict(do_sample=False, max_new_tokens=max_new_tokens, use_cache=False, \n",
|
||||
"# output_hidden_states=True, return_dict=True,\n",
|
||||
"# )\n",
|
||||
@@ -349,7 +356,7 @@
|
||||
"\n",
|
||||
"# for i,s,p,n in zip(inputs, baseline_outputs['text_ans'], control_outputs['text_ans'], control_outputs_neg['text_ans']):\n",
|
||||
"# print(\"===== No Control =====\")\n",
|
||||
"# print(S)\n",
|
||||
"# print(s)\n",
|
||||
"# print(f\"===== + Honesty Control =====\")\n",
|
||||
"# print(p)\n",
|
||||
"# print()\n",
|
||||
@@ -374,9 +381,20 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"execution_count": 12,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<src.repe.rep_control_pipeline_baukit.RepControlPipeline2 at 0x7fcca5807130>"
|
||||
]
|
||||
},
|
||||
"execution_count": 12,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"layer_id = hidden_layers\n",
|
||||
"\n",
|
||||
@@ -391,13 +409,36 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"execution_count": 47,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"===== No Control =====\n",
|
||||
"negative\n",
|
||||
"===== + Honesty Control =====\n",
|
||||
"negative\n",
|
||||
"\n",
|
||||
"===== - Honesty Control =====\n",
|
||||
"negative\n",
|
||||
"\n",
|
||||
"===== No Control =====\n",
|
||||
"The\n",
|
||||
"===== + Honesty Control =====\n",
|
||||
"The\n",
|
||||
"\n",
|
||||
"===== - Honesty Control =====\n",
|
||||
"negative\n",
|
||||
"\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"inputs = dataset_train[:2]\n",
|
||||
"coeff=8.0\n",
|
||||
"max_new_tokens=1\n",
|
||||
"max_new_tokens=3\n",
|
||||
"text_gen_kwargs = dict(do_sample=False, max_new_tokens=max_new_tokens, use_cache=False, \n",
|
||||
" output_hidden_states=True, return_dict=True, max_length=cfg.max_length,\n",
|
||||
" )\n",
|
||||
@@ -416,23 +457,56 @@
|
||||
" control_outputs_neg = rep_control_pipeline(inputs, activations=activations_neg, batch_size=batch_size, **text_gen_kwargs)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"for i,s,p,n in zip(inputs, baseline_outputs, control_outputs, control_outputs_neg):\n",
|
||||
"for i,s,p,n in zip(inputs, baseline_outputs['text_ans'], control_outputs['text_ans'], control_outputs_neg['text_ans']):\n",
|
||||
" print(\"===== No Control =====\")\n",
|
||||
" print(s['generated_text'][0].replace(i, \"\"))\n",
|
||||
" print(s)\n",
|
||||
" print(f\"===== + Honesty Control =====\")\n",
|
||||
" print(p['generated_text'][0].replace(i, \"\"))\n",
|
||||
" print(p)\n",
|
||||
" print()\n",
|
||||
" print(f\"===== - Honesty Control =====\")\n",
|
||||
" print(n['generated_text'][0].replace(i, \"\"))\n",
|
||||
" print(n)\n",
|
||||
" print()\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"execution_count": 71,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"===== No Control =====\n",
|
||||
"2.03%\n",
|
||||
"===== + Honesty Control =====\n",
|
||||
"23.10%\n",
|
||||
"\n",
|
||||
"===== - Honesty Control =====\n",
|
||||
"1.02%\n",
|
||||
"\n",
|
||||
"===== No Control =====\n",
|
||||
"33.63%\n",
|
||||
"===== + Honesty Control =====\n",
|
||||
"83.52%\n",
|
||||
"\n",
|
||||
"===== - Honesty Control =====\n",
|
||||
"7.75%\n",
|
||||
"\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"for i,s,p,n in zip(inputs, baseline_outputs['ans'], control_outputs['ans'], control_outputs_neg['ans']):\n",
|
||||
" print(\"===== No Control =====\")\n",
|
||||
" print(f'{s:02.2%}')\n",
|
||||
" print(f\"===== + Honesty Control =====\")\n",
|
||||
" print(f'{p:02.2%}')\n",
|
||||
" print()\n",
|
||||
" print(f\"===== - Honesty Control =====\")\n",
|
||||
" print(f'{n:02.2%}')\n",
|
||||
" print()\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
from transformers import AutoModel, AutoModelForCausalLM
|
||||
from transformers.pipelines import PIPELINE_REGISTRY
|
||||
from .rep_reading_pipeline import RepReadingPipeline
|
||||
from .rep_control_pipeline import RepControlPipeline, RepControlPipeline2
|
||||
from .rep_control_pipeline import RepControlPipeline
|
||||
from .rep_control_pipeline_baukit import RepControlPipeline2
|
||||
|
||||
def repe_pipeline_registry():
|
||||
PIPELINE_REGISTRY.register_pipeline(
|
||||
|
||||
@@ -7,9 +7,11 @@ from transformers.pipelines import (
|
||||
from transformers.pipelines.base import GenericTensor
|
||||
from .rep_control_reading_vec import WrappedReadingVecModel
|
||||
from typing import Dict
|
||||
from .rep_control_pipeline_baukit import RepControlPipeline2
|
||||
|
||||
|
||||
class RepControlPipeline(FeatureExtractionPipeline):
|
||||
|
||||
class RepControlPipeline(RepControlPipeline2):
|
||||
def __init__(
|
||||
self,
|
||||
model,
|
||||
@@ -35,12 +37,6 @@ class RepControlPipeline(FeatureExtractionPipeline):
|
||||
|
||||
super().__init__(model=model, tokenizer=tokenizer, **kwargs)
|
||||
|
||||
def preprocess(self, inputs, **tokenize_kwargs) -> Dict[str, GenericTensor]:
|
||||
# tokenize a batch of inputs
|
||||
return_tensors = self.framework
|
||||
model_inputs = self.tokenizer(inputs['question'], return_tensors=return_tensors, return_attention_mask=True, add_special_tokens=True, truncation=True, padding="max_length", max_length=self.max_length, **tokenize_kwargs)
|
||||
return {**inputs, **model_inputs}
|
||||
|
||||
def __call__(self, text_inputs, activations=None, **kwargs):
|
||||
if activations is not None:
|
||||
self.wrapped_model.reset()
|
||||
@@ -50,112 +46,3 @@ class RepControlPipeline(FeatureExtractionPipeline):
|
||||
self.wrapped_model.reset()
|
||||
|
||||
return outputs
|
||||
|
||||
def _forward(self, model_inputs):
|
||||
inputs = dict(input_ids=model_inputs['input_ids'], attention_mask=model_inputs['attention_mask'])
|
||||
inputs.update(
|
||||
{"use_cache": False, "output_hidden_states": True, "return_dict": True}
|
||||
)
|
||||
with torch.no_grad():
|
||||
model_outputs = self.model(**inputs)
|
||||
|
||||
# retain some of the inputs
|
||||
keep_cols = ["answer_choices", "input_ids", "attention_mask"]
|
||||
model_outputs = {**model_inputs, **model_outputs}
|
||||
return model_outputs
|
||||
|
||||
def postprocess(self, o):
|
||||
# note this sometimes deals with a batch, sometimes with a single result. infuriating
|
||||
assert isinstance(o, dict) and o['logits'].ndim==3, f"expected dict with logits of shape (batch, seq, vocab), got {o['logits'].shape}"
|
||||
# This is called once for each result, but the text pipeline is set up to hande multiple...
|
||||
# This is called once for each result, but the text pipeline is set up to hande multiple...
|
||||
o["end_logits"] = o["logits"][:, -1, :].float()
|
||||
# hidden_states = list(o.hidden_states)
|
||||
o["input_truncated"] = self.tokenizer.batch_decode(o['input_ids'])
|
||||
o["truncated"] = torch.sum(o["attention_mask"], 1)==self.max_length
|
||||
o["text_ans"] = self.tokenizer.batch_decode(o["end_logits"].argmax(-1))
|
||||
o['choice_ids'] = row_choice_ids(o, self.tokenizer)
|
||||
return o
|
||||
|
||||
|
||||
from typing import List, Tuple, Dict, Any, Union, NewType
|
||||
from baukit.nethook import Trace, TraceDict, recursive_copy
|
||||
from src.datasets.intervene import InterventionDict, intervention_meta_fn
|
||||
from functools import partial
|
||||
from src.datasets.scores import choice2ids
|
||||
|
||||
Activations = NewType("InterventionDict", Dict[str, torch.Tensor])
|
||||
|
||||
def row_choice_ids(answer_choices, tokenizer):
|
||||
return choice2ids([[c] for c in answer_choices], tokenizer)
|
||||
|
||||
def intervention_meta_fn2(
|
||||
output: torch.Tensor, layer_name: str, activations: Activations
|
||||
) -> torch.Tensor:
|
||||
"""see
|
||||
- honest_llama: https://github.com/likenneth/honest_llama/blob/e010f82bfbeaa4326cef8493b0dd5b8b14c6da67/validation/validate_2fold.py#L114
|
||||
- baukit: https://github.com/davidbau/baukit/blob/main/baukit/nethook.py#L42C1-L45C56
|
||||
|
||||
Usage:
|
||||
edit_output = partial(intervention_meta_fn2, activations=activations)
|
||||
with TraceDict(model, layers_to_intervene, edit_output=edit_output) as ret:
|
||||
...
|
||||
|
||||
"""
|
||||
for activation in activations[layer_name]:
|
||||
# TODO might be model specific?
|
||||
output[:, :, :] += torch.from_numpy(activation).to(output.device)[None, None, :]
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class RepControlPipeline2(FeatureExtractionPipeline):
|
||||
"""This version uses baukit."""
|
||||
def __init__(self, model, tokenizer, max_length, **kwargs):
|
||||
super().__init__(model=model, tokenizer=tokenizer, **kwargs)
|
||||
self.max_length = max_length
|
||||
|
||||
def __call__(self, model_inputs, activations=None, **kwargs):
|
||||
if activations is not None:
|
||||
# FIXME model specific
|
||||
layers_names = [f'model.model.layers.{i}.post_attention_layernorm' for i in activations.keys()]
|
||||
edit_fn = partial(intervention_meta_fn2, activations=activations)
|
||||
with TraceDict(
|
||||
self.model, layers_names, detach=True, edit_output=edit_fn
|
||||
) as ret:
|
||||
outputs = super().__call__(model_inputs, **kwargs)
|
||||
else:
|
||||
outputs = super().__call__(model_inputs, **kwargs)
|
||||
return outputs
|
||||
|
||||
def preprocess(self, inputs, **tokenize_kwargs) -> Dict[str, GenericTensor]:
|
||||
# tokenize a batch of inputs
|
||||
return_tensors = self.framework
|
||||
model_inputs = self.tokenizer(inputs['question'], return_tensors=return_tensors, return_attention_mask=True, add_special_tokens=True, truncation=True, padding="max_length", max_length=self.max_length, **tokenize_kwargs)
|
||||
return {**inputs, **model_inputs}
|
||||
|
||||
def _forward(self, model_inputs):
|
||||
inputs = dict(input_ids=model_inputs['input_ids'], attention_mask=model_inputs['attention_mask'])
|
||||
inputs.update(
|
||||
{"use_cache": False, "output_hidden_states": True, "return_dict": True}
|
||||
)
|
||||
with torch.no_grad():
|
||||
model_outputs = self.model(**inputs)
|
||||
|
||||
# retain some of the inputs
|
||||
keep_cols = ["answer_choices", "input_ids", "attention_mask"]
|
||||
model_outputs = {**model_inputs, **model_outputs}
|
||||
return model_outputs
|
||||
|
||||
def postprocess(self, o):
|
||||
# note this sometimes deals with a batch, sometimes with a single result. infuriating
|
||||
assert isinstance(o, dict) and o['logits'].ndim==3, f"expected dict with logits of shape (batch, seq, vocab), got {o['logits'].shape}"
|
||||
# This is called once for each result, but the text pipeline is set up to hande multiple...
|
||||
# This is called once for each result, but the text pipeline is set up to hande multiple...
|
||||
o["end_logits"] = o["logits"][:, -1, :].float()
|
||||
# hidden_states = list(o.hidden_states)
|
||||
o["input_truncated"] = self.tokenizer.batch_decode(o['input_ids'])
|
||||
o["truncated"] = torch.sum(o["attention_mask"], 1)==self.max_length
|
||||
o["text_ans"] = self.tokenizer.batch_decode(o["end_logits"].argmax(-1))
|
||||
o['choice_ids'] = row_choice_ids(o, self.tokenizer)
|
||||
return o
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
import torch
|
||||
from transformers.pipelines import (
|
||||
TextGenerationPipeline,
|
||||
FeatureExtractionPipeline,
|
||||
Pipeline,
|
||||
)
|
||||
from transformers.pipelines.base import GenericTensor
|
||||
from .rep_control_reading_vec import WrappedReadingVecModel
|
||||
from typing import List, Tuple, Dict, Any, Union, NewType
|
||||
from baukit.nethook import Trace, TraceDict, recursive_copy
|
||||
from functools import partial
|
||||
from src.datasets.scores import choice2ids
|
||||
from src.datasets.scores import scores2choice_probs
|
||||
|
||||
|
||||
Activations = NewType("InterventionDict", Dict[str, torch.Tensor])
|
||||
|
||||
def row_choice_ids(answer_choices, tokenizer):
|
||||
return choice2ids([[c] for c in answer_choices], tokenizer)
|
||||
|
||||
|
||||
def intervene(output, activation):
|
||||
assert output.ndim == 3, f"expected output to be (batch, seq, vocab), got {output.shape}"
|
||||
return output + activation.to(output.device)[None, None, :]
|
||||
|
||||
def intervention_meta_fn2(
|
||||
outputs: torch.Tensor, layer_name: str, activations: Activations
|
||||
) -> torch.Tensor:
|
||||
"""see
|
||||
- honest_llama: https://github.com/likenneth/honest_llama/blob/e010f82bfbeaa4326cef8493b0dd5b8b14c6da67/validation/validate_2fold.py#L114
|
||||
- baukit: https://github.com/davidbau/baukit/blob/main/baukit/nethook.py#L42C1-L45C56
|
||||
|
||||
Usage:
|
||||
edit_output = partial(intervention_meta_fn2, activations=activations)
|
||||
with TraceDict(model, layers_to_intervene, edit_output=edit_output) as ret:
|
||||
...
|
||||
|
||||
"""
|
||||
if type(outputs) is tuple:
|
||||
output0 = intervene(outputs[0], activations[layer_name])
|
||||
return tuple(output0, *outputs[1:])
|
||||
elif type(outputs) is torch.Tensor:
|
||||
return intervene(outputs, activations[layer_name])
|
||||
else:
|
||||
raise ValueError(f"outputs must be tuple or tensor, got {type(outputs)}")
|
||||
|
||||
|
||||
class RepControlPipeline2(FeatureExtractionPipeline):
|
||||
"""This version uses baukit."""
|
||||
def __init__(self, model, tokenizer, max_length, layer_name_tmpl="model.layers.{}", **kwargs):
|
||||
super().__init__(model=model, tokenizer=tokenizer, **kwargs)
|
||||
self.max_length = max_length
|
||||
self.layer_name_tmpl = layer_name_tmpl
|
||||
|
||||
def __call__(self, model_inputs, activations=None, **kwargs):
|
||||
if activations is not None:
|
||||
activations_i = Activations({self.layer_name_tmpl.format(k):v for k,v in activations.items()})
|
||||
layers_names = [self.layer_name_tmpl.format(i) for i in activations.keys()]
|
||||
edit_fn = partial(intervention_meta_fn2, activations=activations_i)
|
||||
with TraceDict(
|
||||
self.model, layers_names, detach=True, edit_output=edit_fn
|
||||
) as ret:
|
||||
outputs = super().__call__(model_inputs, **kwargs)
|
||||
else:
|
||||
outputs = super().__call__(model_inputs, **kwargs)
|
||||
return outputs
|
||||
|
||||
def preprocess(self, inputs, **tokenize_kwargs) -> Dict[str, GenericTensor]:
|
||||
# tokenize a batch of inputs
|
||||
return_tensors = self.framework
|
||||
model_inputs = self.tokenizer(inputs['question'], return_tensors=return_tensors, return_attention_mask=True, add_special_tokens=True, truncation=True, padding="max_length", max_length=self.max_length, **tokenize_kwargs)
|
||||
return {**inputs, **model_inputs}
|
||||
|
||||
def _forward(self, model_inputs):
|
||||
inputs = dict(input_ids=model_inputs['input_ids'], attention_mask=model_inputs['attention_mask'])
|
||||
inputs.update(
|
||||
{"use_cache": False, "output_hidden_states": True, "return_dict": True}
|
||||
)
|
||||
with torch.no_grad():
|
||||
model_outputs = self.model(**inputs)
|
||||
|
||||
# retain some of the inputs
|
||||
keep_cols = ["answer_choices", "input_ids", "attention_mask"]
|
||||
model_outputs = {**model_inputs, **model_outputs}
|
||||
return model_outputs
|
||||
|
||||
def postprocess(self, o):
|
||||
# note this sometimes deals with a batch, sometimes with a single result. infuriating
|
||||
assert isinstance(o, dict) and o['logits'].ndim==3, f"expected dict with logits of shape (batch, seq, vocab), got {o['logits'].shape}"
|
||||
# This is called once for each result, but the text pipeline is set up to hande multiple...
|
||||
# This is called once for each result, but the text pipeline is set up to hande multiple...
|
||||
o["end_logits"] = o["logits"][:, -1, :].float()
|
||||
# hidden_states = list(o.hidden_states)
|
||||
o["input_truncated"] = self.tokenizer.batch_decode(o['input_ids'])
|
||||
o["truncated"] = torch.sum(o["attention_mask"], 1)==self.max_length
|
||||
o["text_ans"] = self.tokenizer.batch_decode(o["end_logits"].argmax(-1))
|
||||
o['choice_ids'] = [row_choice_ids(ac, self.tokenizer) for ac in o['answer_choices']]
|
||||
|
||||
p = o['add_ans'] = torch.stack([scores2choice_probs2(l, c) for l, c in zip(o['end_logits'], o['choice_ids'])])
|
||||
o['ans'] = p[:, 1] / (torch.sum(p, 1) + 1e-5)
|
||||
|
||||
return o
|
||||
|
||||
def scores2choice_probs2(logits, choiceids: List[List[int]]):
|
||||
"""calculate the probability for each group of choices."""
|
||||
assert logits.ndim==1, f"expected logits to be 1d, got {logits.shape}"
|
||||
probs = torch.softmax(logits, 0) # shape [tokens, inferences)
|
||||
probs_c = torch.tensor([[probs[cc] for cc in c] for c in choiceids]).sum(1) # sum over alternate choices e.g. [['decrease', 'dec'],['inc', 'increase']]
|
||||
return probs_c
|
||||
Reference in New Issue
Block a user