model_torch.py 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177
  1. import torch.nn as nn
  2. import torch
  3. class TableHeadModel(nn.Module):
  4. def __init__(self):
  5. super(TableHeadModel, self).__init__()
  6. self.char_num = 20
  7. self.char_embed = 60
  8. self.char_embed_expand = 128
  9. self.dense0 = nn.Linear(self.char_embed, self.char_embed_expand)
  10. self.dense3 = nn.Linear(self.char_num * self.char_embed_expand, 64)
  11. self.dense4 = nn.Linear(64, 1)
  12. self.sigmoid = nn.Sigmoid()
  13. self.ln_dnn_2 = nn.LayerNorm([64])
  14. self.device = torch.device("cpu")
  15. self.relu = nn.LeakyReLU()
  16. self.dropout = nn.Dropout(0.3)
  17. self.cnn1d_0 = nn.Conv1d(self.char_embed_expand,
  18. self.char_embed_expand,
  19. (3,), padding=self.get_padding(3))
  20. self.cnn1d_1 = nn.Conv1d(self.char_embed_expand,
  21. self.char_embed_expand,
  22. (3,), padding=self.get_padding(3))
  23. self.cnn3d_0 = nn.Conv3d(self.char_embed_expand, self.char_embed_expand,
  24. (3, 3, 3), padding=self.get_padding(3))
  25. self.cnn3d_1 = nn.Conv3d(self.char_embed_expand, self.char_embed_expand,
  26. (3, 3, 3), padding=self.get_padding(3))
  27. def get_padding(self, kernel_size, stride=1):
  28. return (kernel_size - 1) // 2 * stride
  29. def forward(self, x):
  30. batch, row, col, char_num, char_embed = x.shape
  31. # cnn 1d
  32. cnn1d_x = torch.squeeze(x, 0)
  33. cnn1d_x = cnn1d_x.view([row*col, char_num, char_embed])
  34. cnn1d_x = self.dense0(cnn1d_x)
  35. cnn1d_x = torch.permute(cnn1d_x, [0, 2, 1])
  36. cnn1d_x = self.cnn1d_0(cnn1d_x)
  37. cnn1d_x = self.relu(cnn1d_x)
  38. cnn1d_x = self.dropout(cnn1d_x)
  39. cnn1d_x = self.cnn1d_1(cnn1d_x)
  40. cnn1d_x = self.relu(cnn1d_x)
  41. cnn1d_x = self.dropout(cnn1d_x)
  42. cnn1d_x = torch.permute(cnn1d_x, [0, 2, 1])
  43. cnn1d_x = cnn1d_x.contiguous().view(row, col, char_num, self.char_embed_expand)
  44. cnn1d_x = torch.unsqueeze(cnn1d_x, 0)
  45. # print(cnn1d_x.shape)
  46. # cnn 3d
  47. cnn3d_x = torch.permute(cnn1d_x, [0, 4, 3, 1, 2])
  48. cnn3d_x = self.cnn3d_0(cnn3d_x)
  49. cnn3d_x = self.relu(cnn3d_x)
  50. cnn3d_x = self.dropout(cnn3d_x)
  51. cnn3d_x = self.cnn3d_1(cnn3d_x)
  52. cnn3d_x = self.relu(cnn3d_x)
  53. cnn3d_x = self.dropout(cnn3d_x)
  54. cnn3d_x = torch.squeeze(cnn3d_x, 0)
  55. cnn3d_x = torch.permute(cnn3d_x, [2, 3, 1, 0])
  56. cnn3d_x = cnn3d_x.contiguous().view(row, col, char_num * self.char_embed_expand)
  57. # dnn
  58. x = self.dense3(cnn3d_x)
  59. x = self.ln_dnn_2(x)
  60. x = self.relu(x)
  61. x = self.dense4(x)
  62. x = self.sigmoid(x)
  63. x = torch.squeeze(x, -1)
  64. return x
  65. class TableHeadModel2(nn.Module):
  66. def __init__(self):
  67. super(TableHeadModel2, self).__init__()
  68. self.char_num = 20
  69. self.char_embed = 60
  70. self.char_embed_expand = 128
  71. self.dense0 = nn.Linear(self.char_embed, self.char_embed_expand)
  72. self.dense3 = nn.Linear(self.char_num * self.char_embed_expand, 64)
  73. self.dense4 = nn.Linear(64, 1)
  74. self.sigmoid = nn.Sigmoid()
  75. self.ln_dnn_2 = nn.LayerNorm([64])
  76. self.device = torch.device("cpu")
  77. self.relu = nn.LeakyReLU()
  78. self.dropout = nn.Dropout(0.6)
  79. # self.cnn1d_0 = nn.Conv1d(self.char_embed_expand,
  80. # self.char_embed_expand,
  81. # (3,), padding=self.get_padding(3))
  82. # self.cnn1d_1 = nn.Conv1d(self.char_embed_expand,
  83. # self.char_embed_expand,
  84. # (3,), padding=self.get_padding(3))
  85. encoder_layer1 = nn.TransformerEncoderLayer(d_model=self.char_embed_expand, nhead=2,
  86. dim_feedforward=128, batch_first=True)
  87. self.transformer1 = nn.TransformerEncoder(encoder_layer1, 2)
  88. self.ln_encoder_0 = nn.LayerNorm([self.char_embed_expand])
  89. self.cnn3d_0 = nn.Conv3d(self.char_embed_expand, self.char_embed_expand,
  90. (3, 3, 3), padding=self.get_padding(3))
  91. self.cnn3d_1 = nn.Conv3d(self.char_embed_expand, self.char_embed_expand,
  92. (3, 3, 3), padding=self.get_padding(3))
  93. # self.cnn3d_2 = nn.Conv3d(self.char_embed, self.char_embed,
  94. # (3, 3, 3), padding=self.get_padding(3))
  95. def get_padding(self, kernel_size, stride=1):
  96. return (kernel_size - 1) // 2 * stride
  97. def forward(self, x):
  98. batch, row, col, char_num, char_embed = x.shape
  99. # Embedding
  100. x = torch.squeeze(x, 0)
  101. x = x.view([row*col, char_num, char_embed])
  102. x = self.dense0(x)
  103. # transformer
  104. box_attention = self.transformer1(x)
  105. box_attention = self.ln_encoder_0(box_attention)
  106. box_attention = torch.permute(box_attention, [0, 2, 1])
  107. box_attention = box_attention.contiguous().view(row, col, char_num, self.char_embed_expand)
  108. box_attention = torch.unsqueeze(box_attention, 0)
  109. # cnn1d_x = torch.permute(cnn1d_x, [0, 2, 1])
  110. # cnn1d_x = self.cnn1d_0(cnn1d_x)
  111. # cnn1d_x = self.relu(cnn1d_x)
  112. # cnn1d_x = self.dropout(cnn1d_x)
  113. # cnn1d_x = self.cnn1d_1(cnn1d_x)
  114. # cnn1d_x = self.relu(cnn1d_x)
  115. # cnn1d_x = self.dropout(cnn1d_x)
  116. #
  117. # cnn1d_x = torch.permute(cnn1d_x, [0, 2, 1])
  118. # cnn1d_x = cnn1d_x.contiguous().view(row, col, char_num, self.char_embed_expand)
  119. # cnn1d_x = torch.unsqueeze(cnn1d_x, 0)
  120. # print(cnn1d_x.shape)
  121. # cnn 3d
  122. cnn3d_x = torch.permute(box_attention, [0, 4, 3, 1, 2])
  123. cnn3d_x = self.cnn3d_0(cnn3d_x)
  124. cnn3d_x = self.relu(cnn3d_x)
  125. cnn3d_x = self.dropout(cnn3d_x)
  126. cnn3d_x = self.cnn3d_1(cnn3d_x)
  127. cnn3d_x = self.relu(cnn3d_x)
  128. cnn3d_x = self.dropout(cnn3d_x)
  129. cnn3d_x = torch.squeeze(cnn3d_x, 0)
  130. cnn3d_x = torch.permute(cnn3d_x, [2, 3, 1, 0])
  131. cnn3d_x = cnn3d_x.contiguous().view(row, col, char_num * self.char_embed_expand)
  132. # dnn
  133. x = self.dense3(cnn3d_x)
  134. x = self.ln_dnn_2(x)
  135. x = self.relu(x)
  136. x = self.dense4(x)
  137. x = self.sigmoid(x)
  138. x = torch.squeeze(x, -1)
  139. return x