-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathbert_embedding.py
More file actions
94 lines (80 loc) · 3.34 KB
/
Copy pathbert_embedding.py
File metadata and controls
94 lines (80 loc) · 3.34 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
import json
import random
import copy
import pandas as pd
from transformers import BertTokenizer, BertModel
import torch
from tqdm import tqdm
from q2e_model import BertSimilarity
import argparse
random.seed(1234)
@torch.no_grad()
def bert_embedding(dataset, logs_file: str):
logs = load_logs(logs_file)
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
bert_model = BertModel.from_pretrained('bert-base-uncased')
d = {}
for index, raw in tqdm(logs.iterrows()):
# print(log)
log_input = tokenizer(raw['EventTemplate'], return_tensors="pt")
log_input = {k:v[:, :512] for k, v in log_input.items()} # max input length is 512
log_output = bert_model(**log_input)
log_vec = log_output.last_hidden_state.squeeze()[-1] # used last state to instead of the text
d[raw['EventTemplate']] = log_vec.detach().tolist()
with open('./logs/{}/event2vec.json'.format(dataset), 'w') as f:
json.dump(d, f)
@torch.no_grad()
def my_bert_embedding(dataset, logs_file: str):
logs = load_logs(logs_file)
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
model = BertSimilarity()
model.load_state_dict(torch.load('./logs/{}/mybert.pth'.format(dataset)))
d = {}
for index, raw in tqdm(logs.iterrows()):
# print(log)
log_input = tokenizer(raw['EventTemplate'], max_length=512, padding=True, truncation=True, return_tensors="pt")
log_vec = model.forward_once(log_input)
d[raw['EventTemplate']] = log_vec.squeeze().detach().tolist()
with open('./logs/{}/event2vec_mybert.json'.format(dataset), 'w') as f:
json.dump(d, f)
def load_json(fn:str) -> list:
with open(fn, 'r') as f:
lines = []
for line in f.readlines():
dic = json.loads(line)
lines.append(dic)
return lines
def save_data(data: list) -> None:
pass
def split_data(data: list, train_val_test=(0.6, 0.1, 0.3)) -> dict:
result = {}
random.shuffle(data) # shuffle
total_count = len(data)
train_ratio, val_ratio, test_ratio = train_val_test
train_count = int(total_count * train_ratio)
val_count = int(total_count * val_ratio)
test_count = total_count - train_count - val_count
result['train'] = copy.deepcopy(data[:train_count])
result['val'] = copy.deepcopy(data[train_count: train_count + val_count])
result['test'] = copy.deepcopy(data[train_count + val_count:])
# save
# print('save data...')
# save_data(result['train'])
# save_data(result['val'])
# save_data(result['test'])
print("total: {}, train/val/test: {}/{}/{}".format(total_count, train_count, val_count, test_count))
return result
def load_logs(fn: str):
df = pd.read_csv(fn)
return df
def load_qa(qa_file:str) -> dict:
qa_list = load_json(qa_file)
datasets = split_data(qa_list)
return datasets
if __name__ == '__main__':
argparser = argparse.ArgumentParser()
argparser.add_argument('--dataset', type=str, help='dataset to use')
arg = argparser.parse_args()
dataset = arg.dataset
bert_embedding(dataset, './logs/{}/{}_2k.log_templates.csv'.format(dataset, dataset))
# my_bert_embedding(dataset, './logs/{}/{}_2k.log_templates.csv'.format(dataset, dataset))