| 24 | |
| 25 | |
| 26 | class Teacher_Model(nn.Module): |
| 27 | def __init__(self, n_users, n_items, embedding_dim, weight_size, dropout_list, image_feats, text_feats): |
| 28 | |
| 29 | super().__init__() |
| 30 | self.n_users = n_users |
| 31 | self.n_items = n_items |
| 32 | self.embedding_dim = embedding_dim |
| 33 | self.weight_size = weight_size |
| 34 | self.n_ui_layers = len(self.weight_size) |
| 35 | self.weight_size = [self.embedding_dim] + self.weight_size |
| 36 | |
| 37 | self.image_trans = nn.Linear(image_feats.shape[1], args.embed_size) |
| 38 | self.text_trans = nn.Linear(text_feats.shape[1], args.embed_size) |
| 39 | nn.init.xavier_uniform_(self.image_trans.weight) |
| 40 | nn.init.xavier_uniform_(self.text_trans.weight) |
| 41 | self.encoder = nn.ModuleDict() |
| 42 | self.encoder['image_encoder'] = self.image_trans # ^-^ |
| 43 | self.encoder['text_encoder'] = self.text_trans # ^-^ |
| 44 | |
| 45 | self.user_id_embedding = nn.Embedding(n_users, self.embedding_dim) |
| 46 | self.item_id_embedding = nn.Embedding(n_items, self.embedding_dim) |
| 47 | |
| 48 | nn.init.xavier_uniform_(self.user_id_embedding.weight) |
| 49 | nn.init.xavier_uniform_(self.item_id_embedding.weight) |
| 50 | self.image_feats = torch.tensor(image_feats).float().cuda() |
| 51 | self.text_feats = torch.tensor(text_feats).float().cuda() |
| 52 | self.image_embedding = nn.Embedding.from_pretrained(torch.Tensor(image_feats), freeze=False) |
| 53 | self.text_embedding = nn.Embedding.from_pretrained(torch.Tensor(text_feats), freeze=False) |
| 54 | |
| 55 | self.softmax = nn.Softmax(dim=-1) |
| 56 | self.act = nn.Sigmoid() |
| 57 | self.sigmoid = nn.Sigmoid() |
| 58 | self.dropout = nn.Dropout(p=args.drop_rate) |
| 59 | self.batch_norm = nn.BatchNorm1d(args.embed_size) |
| 60 | |
| 61 | def mm(self, x, y): |
| 62 | if args.sparse: |
| 63 | return torch.sparse.mm(x, y) |
| 64 | else: |
| 65 | return torch.mm(x, y) |
| 66 | def sim(self, z1, z2): |
| 67 | z1 = F.normalize(z1) |
| 68 | z2 = F.normalize(z2) |
| 69 | return torch.mm(z1, z2.t()) |
| 70 | |
| 71 | def batched_contrastive_loss(self, z1, z2, batch_size=4096): |
| 72 | device = z1.device |
| 73 | num_nodes = z1.size(0) |
| 74 | num_batches = (num_nodes - 1) // batch_size + 1 |
| 75 | f = lambda x: torch.exp(x / self.tau) |
| 76 | indices = torch.arange(0, num_nodes).to(device) |
| 77 | losses = [] |
| 78 | |
| 79 | for i in range(num_batches): |
| 80 | mask = indices[i * batch_size:(i + 1) * batch_size] |
| 81 | refl_sim = f(self.sim(z1[mask], z1)) |
| 82 | between_sim = f(self.sim(z1[mask], z2)) |
| 83 | |