pre_process_torch.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132
  1. #coding=utf-8
  2. import os
  3. import sys
  4. import numpy as np
  5. import torch
  6. from torch.utils.data import Dataset
  7. sys.path.append(os.path.abspath(os.path.dirname(__file__) + "/../../../"))
  8. from BiddingKG.dl.common.Utils import embedding_word, embedding_word_forward
  9. def set_label(row, row_label):
  10. if len(row) == 1:
  11. row_label = [0 for x in row]
  12. elif len(set(row)) == 1:
  13. row_label = [0 for x in row]
  14. else:
  15. row_label = [0 if x in ["", " ", "/", '无', '-', '~~'] else row_label[i] for i, x in enumerate(row)]
  16. return row_label
  17. def set_same_table_head(inputs, y_pred1):
  18. inputs = torch.squeeze(inputs, 0)
  19. for i in range(inputs.shape[0]):
  20. for j in range(inputs.shape[1]-1):
  21. col1 = inputs[i, j, :, :]
  22. col2 = inputs[i, j+1, :, :]
  23. if (torch.abs(col1 - col2) < 1e-4).all():
  24. # print('same value', col1[abs(col1) > 0.], col2[abs(col1) > 0.])
  25. if (y_pred1[i, j] <= 0.5 and y_pred1[i, j+1] <= 0.5) or (y_pred1[i, j] > 0.5 and y_pred1[i, j+1] > 0.5):
  26. continue
  27. else:
  28. # print('differ label', y_pred[i, j], y_pred[i, j+1])
  29. y_pred1[i, j+1] = y_pred1[i, j]
  30. for i in range(inputs.shape[1]):
  31. for j in range(inputs.shape[0]-1):
  32. row1 = inputs[j, i, :, :]
  33. row2 = inputs[j+1, i, :, :]
  34. if (torch.abs(row1 - row2) < 1e-4).all():
  35. if (y_pred1[j, i] <= 0.5 and y_pred1[j+1, i] <= 0.5) or (y_pred1[j, i] > 0.5 and y_pred1[j+1, i] > 0.5):
  36. continue
  37. else:
  38. # print('same value', row1[abs(row1) > 0.], row2[abs(row2) > 0.])
  39. # print('differ label', y_pred[i, j], y_pred[i, j+1])
  40. # print('before', x11[0, j, i], x11[0, j+1, i])
  41. y_pred1[j+1, i] = y_pred1[j, i]
  42. # print('after', x1[0, j, i], x1[0, j+1, i])
  43. return y_pred1
  44. def data_to_numpy29(data_list, data_label_list):
  45. """
  46. 输出表格 (table_cnt, row, col, 20, 60)
  47. :param data_list:
  48. :param data_label_list:
  49. :return:
  50. """
  51. data_num = len(data_list)
  52. new_data_list = []
  53. new_label_list = []
  54. mask_list = []
  55. for i in range(len(data_list)):
  56. table = data_list[i]
  57. table_label = []
  58. if data_label_list:
  59. table_label = data_label_list[i]
  60. embed_list = []
  61. label_list = []
  62. mask = []
  63. for j in range(len(table)):
  64. row = table[j]
  65. blank_list = [0 if x in ["", " ", "/"] else 1 for x in row]
  66. mask.append(blank_list)
  67. row = embedding_word_forward(row, shape=(len(row), 20, 60))
  68. embed_list.append(row)
  69. if data_label_list:
  70. row_label = table_label[j]
  71. # print(j, row_label)
  72. row_label = [int(x) for x in row_label]
  73. row_label = set_label(table[j], row_label)
  74. label_list.append(row_label)
  75. embed_list = np.array(embed_list, dtype=np.float32)
  76. label_list = np.array(label_list, dtype=np.float32)
  77. mask = np.array(mask, dtype=np.float32)
  78. # print('embed_list.shape', embed_list.shape)
  79. # print('label_list.shape', label_list.shape)
  80. new_data_list.append(embed_list)
  81. new_label_list.append(label_list)
  82. mask_list.append(mask)
  83. new_data_list = np.array(new_data_list, dtype=np.float32)
  84. new_label_list = np.array(new_label_list, dtype=np.float32)
  85. mask_list = np.array(mask_list, dtype=np.float32)
  86. # print(new_data_list.shape)
  87. return new_data_list, new_label_list, mask_list
  88. class CustomDatasetTiny40(Dataset):
  89. def __init__(self, data_x, data_y, mode=0):
  90. if mode in [0, 1]:
  91. # Split -> Train, Test
  92. split_size = int(len(data_x)*0.1)
  93. test_x, test_y = data_x[:split_size], data_y[:split_size]
  94. train_x, train_y = data_x[split_size:], data_y[split_size:]
  95. if mode == 0:
  96. self.data = train_x
  97. self.targets = train_y
  98. else:
  99. self.data = test_x
  100. self.targets = test_y
  101. else:
  102. pass
  103. # self.data = data
  104. # self.targets = targets
  105. def __len__(self):
  106. return len(self.data)
  107. def __getitem__(self, idx):
  108. # x, y = data_to_numpy12([self.data[idx]], [self.targets[idx]])
  109. x, y, mask = data_to_numpy29([self.data[idx]], [self.targets[idx]])
  110. x = x[0]
  111. y = y[0]
  112. mask = mask[0]
  113. return x, y, mask