mirror of
https://github.com/wassname/discovering_latent_knowledge.git
synced 2026-09-10 12:00:13 +08:00
fixed bug in batch
This commit is contained in:
@@ -2,7 +2,7 @@
|
||||
from tqdm.auto import tqdm
|
||||
from src.datasets.hs import ExtractHiddenStates
|
||||
from torch.utils.data import DataLoader
|
||||
from datasets import Dataset
|
||||
from datasets.arrow_dataset import Dataset
|
||||
import hashlib
|
||||
import pickle
|
||||
import numpy as np
|
||||
@@ -23,7 +23,7 @@ def batch_hidden_states(model, tokenizer, data: Dataset, n=100, batch_size=2, mc
|
||||
ds_p_subset = data.select(range(n))
|
||||
ds_p_subset.set_format(type="pandas", columns=['lie', 'label', 'prompt', 'prompt_truncated'])
|
||||
|
||||
dl = DataLoader(ds_t_subset, batch_size=batch_size, shuffle=True)
|
||||
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"]
|
||||
nn = len(input_ids)
|
||||
@@ -47,7 +47,9 @@ 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]
|
||||
info = ds_p_subset[k].iloc[0].to_dict()
|
||||
|
||||
assert info['label']==true_labels[j].item(), 'these should line up'
|
||||
|
||||
yield dict(
|
||||
hs0=hs0['hidden_states'][j],
|
||||
@@ -56,8 +58,8 @@ def batch_hidden_states(model, tokenizer, data: Dataset, n=100, batch_size=2, mc
|
||||
hs1=hs1['hidden_states'][j],
|
||||
scores1=hs1["scores"][j],
|
||||
|
||||
true=true_labels[j].item(),
|
||||
index=index[j],
|
||||
label_b=true_labels[j].item(),
|
||||
ds_index=index[j],
|
||||
|
||||
**info
|
||||
)
|
||||
|
||||
+24
-33
@@ -2,18 +2,20 @@ import torch
|
||||
import torch.nn as nn
|
||||
import lightning as pl
|
||||
import pandas as pd
|
||||
from torch.utils.data import Dataset, DataLoader, TensorDataset
|
||||
from torch.utils.data import DataLoader, TensorDataset
|
||||
from src.datasets.load import ds2df
|
||||
from datasets.arrow_dataset import Dataset
|
||||
|
||||
def make_y(df):
|
||||
# label: is ans2 more true than ans1
|
||||
# so we ask does ans2 have greater probability on "positive" than ans1
|
||||
# then, when the right answer is negative we swap the sign
|
||||
true_switch_sign = df.label*2-1
|
||||
def compute_distance(df):
|
||||
"""distance between ans1 and ans2."""
|
||||
true_switch_sign = df.true*2-1 # switch sign to desired answer. with this we ask which is more true
|
||||
# otherwise we ask which is more positive
|
||||
distance = (df.ans1-df.ans0) * true_switch_sign
|
||||
# y = bool2switch(distance>0)
|
||||
return distance
|
||||
|
||||
to_tensor = lambda x: torch.from_numpy(x).float()
|
||||
to_ds = lambda hs0, hs1, y: TensorDataset(to_tensor(hs0), to_tensor(hs1), to_tensor(y))
|
||||
|
||||
class imdbHSDataModule(pl.LightningDataModule):
|
||||
|
||||
def __init__(self,
|
||||
@@ -22,7 +24,7 @@ class imdbHSDataModule(pl.LightningDataModule):
|
||||
):
|
||||
super().__init__()
|
||||
self.save_hyperparameters(ignore=["ds"])
|
||||
self.ds = ds.shuffle(seed=42)
|
||||
self.ds = ds#.shuffle(seed=42)
|
||||
|
||||
def setup(self, stage: str):
|
||||
h = self.hparams
|
||||
@@ -34,46 +36,35 @@ class imdbHSDataModule(pl.LightningDataModule):
|
||||
)
|
||||
self.df = ds2df(self.ds)
|
||||
|
||||
y_cls = make_y(self.df)
|
||||
y_cls = compute_distance(self.df)
|
||||
|
||||
self.y = y_cls.values
|
||||
self.df['y'] = y_cls
|
||||
|
||||
b = len(self.ds_hs)
|
||||
self.hs1 = self.ds_hs['hs0'].transpose(0, 2, 1)
|
||||
self.hs2 = self.ds_hs['hs1'].transpose(0, 2, 1)
|
||||
self.hs0 = self.ds_hs['hs0'].transpose(0, 2, 1)
|
||||
self.hs1 = self.ds_hs['hs1'].transpose(0, 2, 1)
|
||||
self.ans0 = self.df['ans0'].values
|
||||
self.ans1 = self.df['ans1'].values
|
||||
|
||||
# let's create a simple 50/50 train split (the data is already randomized)
|
||||
n = len(self.y)
|
||||
self.splits = {
|
||||
'train': (0, int(n * 0.5)),
|
||||
'val': (int(n * 0.5), int(n * 0.75)),
|
||||
'test': (int(n * 0.75), n),
|
||||
}
|
||||
|
||||
self.val_split = vs = int(n * 0.5)
|
||||
self.test_split = ts = int(n * 0.75)
|
||||
hs1_train, hs2_train, y_train = self.hs1[:vs], self.hs2[:vs], self.y[:vs]
|
||||
hs1_val, hs2_val, y_val = self.hs1[vs:ts], self.hs2[vs:ts], self.y[vs:ts]
|
||||
hs1_test, hs2_test, y_test = self.hs1[ts:],self. hs2[ts:], self.y[ts:]
|
||||
|
||||
|
||||
to_ds = lambda x0, x1, y: TensorDataset(torch.from_numpy(x0).float(),
|
||||
torch.from_numpy(x1).float(),
|
||||
torch.from_numpy(y).float()
|
||||
)
|
||||
self.datasets = {key: to_ds(self.hs0[start:end], self.hs1[start:end], self.y[start:end]) for key, (start, end) in self.splits.items()}
|
||||
|
||||
self.ds_train = to_ds(hs1_train, hs2_train, y_train)
|
||||
|
||||
self.ds_val = to_ds(hs1_val, hs2_val, y_val)
|
||||
|
||||
self.ds_test = to_ds(hs1_test, hs2_test, y_test)
|
||||
def create_dataloader(self, ds, shuffle=False):
|
||||
return DataLoader(ds, batch_size=self.hparams.batch_size, drop_last=True, shuffle=shuffle)
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(self.ds_train,
|
||||
batch_size=self.hparams.batch_size,
|
||||
drop_last=True,
|
||||
shuffle=True)
|
||||
return self.create_dataloader(self.datasets['train'], shuffle=True)
|
||||
|
||||
def val_dataloader(self):
|
||||
return DataLoader(self.ds_val, batch_size=self.hparams.batch_size, drop_last=True,)
|
||||
return self.create_dataloader(self.datasets['val'])
|
||||
|
||||
def test_dataloader(self):
|
||||
return DataLoader(self.ds_test, batch_size=self.hparams.batch_size, drop_last=True,)
|
||||
return self.create_dataloader(self.datasets['test'])
|
||||
|
||||
+3
-2
@@ -62,7 +62,7 @@ def get_choices_as_tokens(
|
||||
tokenizer, choices:List[str] = ["Positive"], whitespace_first=True
|
||||
) -> List[int]:
|
||||
|
||||
# Note some tokenizers differentiate between "no", "\nno", so we sometime need to add whitespace beforehand...
|
||||
# Note some tokenizers differentiate between "yes", "\nyes" and " yes", so we sometime need to add whitespace beforehand...
|
||||
if not whitespace_first:
|
||||
raise NotImplementedError('TODO')
|
||||
|
||||
@@ -72,7 +72,7 @@ def get_choices_as_tokens(
|
||||
ids.append(id_)
|
||||
|
||||
c2 = tokenizer.decode([id_])
|
||||
assert tokenizer.decode([id_]) == c, f'tokenizer.decode(tokenizer(`{c}`))==`{c2}`!=`{c}`'
|
||||
assert tokenizer.decode([id_]) == c, f'We should be able to encode and decode the choices, but it failed: tokenizer.decode(tokenizer(`{c}`))==`{c2}`!=`{c}`'
|
||||
|
||||
return ids
|
||||
|
||||
@@ -154,6 +154,7 @@ class ExtractHiddenStates:
|
||||
hidden_states=hidden_states,
|
||||
scores=outputs["scores"],
|
||||
input_ids=input_ids,
|
||||
layers=layers,
|
||||
)
|
||||
out = {k: to_numpy(v) for k, v in out.items()}
|
||||
if debug:
|
||||
|
||||
@@ -2,8 +2,10 @@ from pytorch_optimizer import Ranger21
|
||||
import torchmetrics
|
||||
import lightning.pytorch as pl
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torch.nn as nn
|
||||
from torchmetrics import Metric, MetricCollection, Accuracy, AUROC
|
||||
from torchmetrics.functional import accuracy
|
||||
|
||||
from src.helpers import switch2bool, bool2switch
|
||||
|
||||
@@ -15,20 +17,9 @@ class PLRanking(pl.LightningModule):
|
||||
"""
|
||||
def __init__(self, c_in, total_steps, depth=1, hs=16, lr=4e-3, weight_decay=1e-9, dropout=0):
|
||||
super().__init__()
|
||||
# self.probe = MLPProbe(c_in, depth=depth, dropout=dropout, hs=hs)
|
||||
self.probe = None # subclasses must add this
|
||||
self.save_hyperparameters()
|
||||
|
||||
self.loss_fn = nn.SmoothL1Loss()
|
||||
|
||||
# metrics for each stage
|
||||
metrics_template = MetricCollection({
|
||||
'acc': Accuracy(task="binary"),
|
||||
'auroc': AUROC(task="binary")
|
||||
})
|
||||
self.metrics = torch.nn.ModuleDict({
|
||||
f'metrics_{stage}': metrics_template.clone(prefix=stage+'/') for stage in ['train', 'val', 'test']
|
||||
})
|
||||
|
||||
def forward(self, x):
|
||||
return self.probe(x).squeeze(1)
|
||||
|
||||
@@ -40,14 +31,14 @@ class PLRanking(pl.LightningModule):
|
||||
if stage=='pred':
|
||||
return (ypred1-ypred0).float()
|
||||
|
||||
loss = self.loss_fn(ypred1-ypred0, y)
|
||||
self.log(f"{stage}/loss", loss)
|
||||
|
||||
m = self.metrics[f'metrics_{stage}']
|
||||
loss = F.smooth_l1_loss(ypred1-ypred0, y)
|
||||
# self.log(f"{stage}/loss", loss)
|
||||
|
||||
y_cls = switch2bool(ypred1-ypred0)
|
||||
m(y_cls, y>0.)
|
||||
self.log_dict(m, on_epoch=True, on_step=False)
|
||||
self.log_dict({
|
||||
f"{stage}/acc": accuracy(y_cls, y>0, "binary"),
|
||||
f"{stage}/loss": loss,
|
||||
}, on_epoch=True, on_step=False),
|
||||
return loss
|
||||
|
||||
def training_step(self, batch, batch_idx=0, dataloader_idx=0):
|
||||
|
||||
Reference in New Issue
Block a user