mirror of
https://github.com/wassname/discovering_latent_knowledge.git
synced 2026-10-05 12:40:11 +08:00
new dataloading script alpha
This commit is contained in:
1 parent
4cc9daa76d
commit
28dec05dd9
13 files changed
+438
-911
No files matched your search
+40
-39
@@ -1,5 +1,6 @@
|
||||
|
||||
from tqdm.auto import tqdm
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from datasets.arrow_dataset import Dataset
|
||||
import hashlib
|
||||
@@ -8,9 +9,10 @@ import numpy as np
|
||||
|
||||
from src.datasets.hs import ExtractHiddenStates
|
||||
from src.helpers.typing import float_to_int16, int16_to_float
|
||||
from src.helpers.ds import ds_keep_cols
|
||||
|
||||
|
||||
def batch_hidden_states(model, tokenizer, data: Dataset, n=100, batch_size=2, mcdropout=True):
|
||||
def batch_hidden_states(model, tokenizer, data: Dataset, batch_size=2, mcdropout=True):
|
||||
"""
|
||||
Given an encoder-decoder model, a list of data, computes the contrast hidden states on n random examples.
|
||||
Returns numpy arrays of shape (n, hidden_dim) for each candidate label, along with a boolean numpy array of shape (n,)
|
||||
@@ -20,15 +22,16 @@ def batch_hidden_states(model, tokenizer, data: Dataset, n=100, batch_size=2, mc
|
||||
"""
|
||||
ehs = ExtractHiddenStates(model, tokenizer)
|
||||
|
||||
ds_t_subset = data.select(range(n))
|
||||
ds_t_subset.set_format(type='torch', columns=['input_ids', 'label', 'attention_mask'])
|
||||
torch_cols = ['input_ids', 'attention_mask']
|
||||
ds_t_subset = ds_keep_cols(data, torch_cols)
|
||||
ds_t_subset.set_format(type='torch')
|
||||
|
||||
ds_p_subset = data.select(range(n))
|
||||
ds_p_subset.set_format(type="pandas", columns=['lie', 'label', 'prompt', 'prompt_truncated'])
|
||||
ds_p_subset = data.remove_columns(torch_cols)
|
||||
# TODO check it has a few critical ones in
|
||||
|
||||
dl = DataLoader(ds_t_subset, batch_size=batch_size, shuffle=False)
|
||||
for i, batch in enumerate(tqdm(dl, desc='get hidden states')):
|
||||
input_ids, true_labels, attention_mask = batch["input_ids"], batch["label"], batch["attention_mask"]
|
||||
input_ids, attention_mask = batch["input_ids"], batch["attention_mask"]
|
||||
nn = len(input_ids)
|
||||
index = i*batch_size+np.arange(nn)
|
||||
|
||||
@@ -50,57 +53,55 @@ def batch_hidden_states(model, tokenizer, data: Dataset, n=100, batch_size=2, mc
|
||||
for j in range(nn):
|
||||
# let's add the non torch metadata like label, prompt, lie, etc
|
||||
k = i*batch_size + j
|
||||
info = ds_p_subset[k].iloc[0].to_dict()
|
||||
|
||||
assert info['label']==true_labels[j].item(), 'these should line up'
|
||||
info = ds_p_subset[k]
|
||||
|
||||
yield dict(
|
||||
hs0=float_to_int16(hs0['hidden_states'][j]),
|
||||
# int16 makes our storage much smaller
|
||||
hs0=float_to_int16(torch.from_numpy(hs0['hidden_states'][j])),
|
||||
scores0=hs0["scores"][j],
|
||||
|
||||
hs1=float_to_int16(hs1['hidden_states'][j]),
|
||||
hs1=float_to_int16(torch.from_numpy(hs1['hidden_states'][j])),
|
||||
scores1=hs1["scores"][j],
|
||||
|
||||
label_b=true_labels[j].item(),
|
||||
ds_index=index[j],
|
||||
|
||||
**info
|
||||
)
|
||||
|
||||
|
||||
def md5hash(s: bytes) -> str:
|
||||
return hashlib.md5(s).hexdigest()
|
||||
# def md5hash(s: bytes) -> str:
|
||||
# return hashlib.md5(s).hexdigest()
|
||||
|
||||
# unique hash
|
||||
def get_unique_config_hash(prompt_fn, model, tokenizer, data, N):
|
||||
"""
|
||||
generates a unique name
|
||||
# # unique hash
|
||||
# def get_unique_config_hash(cfg, ds_name, split_type):
|
||||
# """
|
||||
# generates a unique name
|
||||
|
||||
datasets would do this use the generation kwargs but this way we have control and can handle non-picklable models and thing like the output of prompt functions if they change
|
||||
# datasets would do this use the generation kwargs but this way we have control and can handle non-picklable models and thing like the output of prompt functions if they change
|
||||
|
||||
# """
|
||||
example_prompt1 = prompt_fn("text", response=0, lie=True)
|
||||
model_repo = model.config._name_or_path
|
||||
# # """
|
||||
# example_prompt1 = prompt_fn("text", response=0, lie=True)
|
||||
# model_repo = model.config._name_or_path
|
||||
|
||||
kwargs = [str(model), str(tokenizer), str(data), str(prompt_fn.__name__), N]
|
||||
key = pickle.dumps(kwargs, 1)
|
||||
hsh = md5hash(key)[:6]
|
||||
# kwargs = [str(model), str(tokenizer), str(data), str(prompt_fn.__name__), N]
|
||||
# key = pickle.dumps(kwargs, 1)
|
||||
# hsh = md5hash(key)[:6]
|
||||
|
||||
sanitize = lambda s:s.replace('/', '').replace('-', '_') if s is not None else s
|
||||
# config_name = f"{sanitize(model_repo)}-N_{N}-ns-{hsh}"
|
||||
# sanitize = lambda s:s.replace('/', '').replace('-', '_') if s is not None else s
|
||||
# # config_name = f"{sanitize(model_repo)}-N_{N}-ns-{hsh}"
|
||||
|
||||
info_kwargs = dict(model_repo=model_repo, config=model.config, data=str(data), prompt_fn=str(prompt_fn.__name__), N=N,
|
||||
example_prompt1=example_prompt1,
|
||||
hsh=hsh)
|
||||
# info_kwargs = dict(model_repo=model_repo, config=model.config, data=str(data), prompt_fn=str(prompt_fn.__name__), N=N,
|
||||
# example_prompt1=example_prompt1,
|
||||
# hsh=hsh)
|
||||
|
||||
return hsh, info_kwargs
|
||||
# return hsh, info_kwargs
|
||||
|
||||
sanitize = lambda s:s.replace('/', '').replace('_', '-') if s is not None else s
|
||||
# sanitize = lambda s:s.replace('/', '').replace('_', '-') if s is not None else s
|
||||
|
||||
def ds_params2fname(dataset_params: dict) -> str:
|
||||
prompt = sanitize(dataset_params['prompt_fmt'].__name__)
|
||||
model_repo = sanitize(dataset_params['model_repo'].split('/')[-1])
|
||||
dataset_name = sanitize(dataset_params['dataset_name'])
|
||||
N = dataset_params['N']
|
||||
N_SHOTS = dataset_params['N_SHOTS']
|
||||
return f"model-{model_repo}_ds-{dataset_name}_{prompt}_N{N}_{N_SHOTS}shots_"
|
||||
# def ds_params2fname(dataset_params: dict) -> str:
|
||||
# prompt = sanitize(dataset_params['prompt_fmt'].__name__)
|
||||
# model_repo = sanitize(dataset_params['model_repo'].split('/')[-1])
|
||||
# dataset_name = sanitize(dataset_params['dataset_name'])
|
||||
# N = dataset_params['N']
|
||||
# N_SHOTS = dataset_params['N_SHOTS']
|
||||
# return f"model-{model_repo}_ds-{dataset_name}_{prompt}_N{N}_{N_SHOTS}shots_"
|
||||
+4
-4
@@ -36,7 +36,7 @@ def label_to_choice(label: bool, class2choices=default_class2choices) -> str:
|
||||
choices = class2choices_to_choices(class2choices)
|
||||
return choices[label]
|
||||
|
||||
def scores2choice_probs(row, class2_ids, keys=["scores0", "scores1"] ):
|
||||
def scores2choice_probs(row, class2_ids: List[int], keys=["scores0", "scores1"] ):
|
||||
""" Given next_token scores (logits) we take only the subset the corresponds to our
|
||||
- negative tokens (e.g. False, no, ...)
|
||||
- and positive tokens (e.g. Yes, yes, affirmative, ...).
|
||||
@@ -52,7 +52,7 @@ def scores2choice_probs(row, class2_ids, keys=["scores0", "scores1"] ):
|
||||
for key in keys:
|
||||
scores = row[key]
|
||||
probs = F.softmax(torch.from_numpy(scores), -1).numpy()
|
||||
probs_c = [probs[class2_ids[c]].sum() for c in class2_ids]
|
||||
probs_c = [probs[c].sum() for c in class2_ids]
|
||||
|
||||
# balance of probs
|
||||
out[key.replace("scores", "choice_probs")] = probs_c
|
||||
@@ -63,8 +63,8 @@ def scores2choice_probs(row, class2_ids, keys=["scores0", "scores1"] ):
|
||||
# out[key.replace("scores", "ansb")] = torch.tensor(scores_c).softmax(-1)[1].item()
|
||||
return out
|
||||
|
||||
def choice2ids(tokenizer, class2hoices: Dict[bool, List[str]]) -> Dict[int, List[int]]:
|
||||
return {k: get_choices_as_tokens(tokenizer, v) for k,v in class2hoices.items()}
|
||||
def choice2ids(tokenizer, class2hoices: List[str]) -> List[int]:
|
||||
return [get_choices_as_tokens(tokenizer, v) for v in class2hoices]
|
||||
|
||||
def get_choices_as_tokens(
|
||||
tokenizer, choices:List[str] = ["Positive"], whitespace_first=True
|
||||
|
||||
Reference in new issue
Block a user