MCPcopy Create free account
hub / github.com/HKUDS/PromptMM / Teacher_Model

Class Teacher_Model

codes/Models.py:26–240  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24
25
26class 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

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected