BBC6521_Project / deepxml /data_utils.py
chen shangheng
deepxml
63239ac
Raw
History Blame Contribute Delete
5.32 kB
import os
import numpy as np
import torch
import joblib
from collections import Counter
from sklearn.preprocessing import MultiLabelBinarizer, normalize
from sklearn.datasets import load_svmlight_file
from tqdm import tqdm
from typing import Union, Iterable
def get_word_emb(vec_path, vocab_path=None):
if vocab_path is not None:
with open(vocab_path) as fp:
vocab = {word: idx for idx, word in enumerate(fp)}
return np.load(vec_path), vocab
else:
return np.load(vec_path)
def get_data(text_file, label_file=None):
return np.load(text_file, allow_pickle=True), np.load(label_file,
allow_pickle=True) if label_file is not None else None
def convert_to_binary(text, max_len=None, vocab=None, pad='<PAD>',
unknown='<UNK>'):
texts = np.asarray([[vocab.get(word, vocab[unknown]) for word in text.split()]], dtype=object)
return truncate_text(texts, max_len, vocab[pad], vocab[unknown])
def truncate_text(texts, max_len=500, padding_idx=0, unknown_idx=1):
if max_len is None:
return texts
texts = np.asarray([list(x[:max_len]) + [padding_idx] * (max_len - len(x)) for x in texts])
texts[(texts == padding_idx).all(axis=1), 0] = unknown_idx
return texts
def get_mlb(mlb_path, labels=None) -> MultiLabelBinarizer:
if os.path.exists(mlb_path):
return joblib.load(mlb_path)
mlb = MultiLabelBinarizer(sparse_output=True)
mlb.fit(labels)
joblib.dump(mlb, mlb_path)
return mlb
def get_sparse_feature(feature_file, label_file):
sparse_x, _ = load_svmlight_file(feature_file, multilabel=True)
return normalize(sparse_x), np.load(label_file) if label_file is not None else None
def load_hierachy(mlb, data_cnf, model_cnf):
label_vocab = {word: i for i, word in enumerate(np.load(data_cnf['label_vocab']))}
edges = set()
hierarchy = set()
label_dict = {}
classes = mlb.classes_.tolist()
num_class = len(mlb.classes_)
with open(data_cnf['hierarchy']) as fin:
for line in fin:
data = line.strip().split()
p = "LABEL_" + str(data[0])
if p in label_vocab.keys():
p = label_vocab[p]
if p not in classes:
continue
p_id = classes.index(p)
for c in data[1:]:
c = "LABEL_" + str(c)
if c in label_vocab.keys():
c = label_vocab[c]
if c not in classes:
continue
c_id = classes.index(c)
edges.add((p, c))
hierarchy.add((p_id, c_id))
with open(data_cnf['labels_dict']) as fin:
for line in fin:
data = line.strip().split()
p = "LABEL_" + str(data[0])
if p in label_vocab.keys():
p = label_vocab[p]
if p in classes:
label_dict[p] = data[1]
def get_root(n):
ret = []
while path_dict[n] != n:
ret.append(n)
n = path_dict[n]
ret.append(n)
return ret
path_dict = {}
for parent, child in edges:
path_dict[child] = parent
for i in classes:
if i not in path_dict:
path_dict[i] = i
inverse_label_list = {}
label_level = {}
# node_list = {}
for i in classes:
inverse_label_list.update({i: get_root(i) + [1]})
label_level.update({i: len(get_root(i))})
def get_distance(node1, node2):
if (node1 not in inverse_label_list) or (node2 not in inverse_label_list):
return 0
node1_lst = inverse_label_list[node1]
node2_lst = inverse_label_list[node2]
p = len(node1_lst) - 1
q = len(node2_lst) - 1
while p > 0 and q > 0 and node1_lst[p] == node2_lst[q]:
p -= 1
q -= 1
return p + q
# Distance Mat
dist_mat_path = data_cnf["dist_mat_path"]
if os.path.exists(dist_mat_path):
distance_matrix = np.load(dist_mat_path)
else:
inputs_label_lst = torch.tensor(classes)
distance_matrix = inputs_label_lst.reshape(1, -1).repeat(inputs_label_lst.size(0), 1)
hier_mat_t = inputs_label_lst.reshape(-1, 1).repeat(1, inputs_label_lst.size(0))
distance_matrix.map_(hier_mat_t, get_distance)
np.save(dist_mat_path, distance_matrix)
# dist_mat_mask = []
# for i in range(1, model_cnf["model"]["dist_max_len"]):
# dist_mat_mask.append((distance_matrix == i).sum(1))
# Edge Mat
edge_mat_path = data_cnf["edge_mat_path"]
if os.path.exists(edge_mat_path):
edge_matrix = np.load(edge_mat_path)
else:
edge_matrix = truncate_text(inverse_label_list.values(), 5)
np.save(edge_mat_path, edge_matrix)
graph_hierarchy = {}
graph_hierarchy["inverse_label_list"] = inverse_label_list # inverse list
graph_hierarchy["label_num"] = num_class # label总数量
graph_hierarchy["hierarchy"] = hierarchy
graph_hierarchy["classes"] = classes
graph_hierarchy["distance_matrix"] = distance_matrix
graph_hierarchy["edge_matrix"] = edge_matrix
return graph_hierarchy