(self, embed_dim, n_intents, temp=0.2, use_multi_head=True, n_heads=4)
| 4 | |
| 5 | class IntentModule(nn.Module): |
| 6 | def __init__(self, embed_dim, n_intents, temp=0.2, use_multi_head=True, n_heads=4): |
| 7 | super(IntentModule, self).__init__() |
| 8 | self.embed_dim = embed_dim |
| 9 | self.n_intents = n_intents |
| 10 | self.temp = temp |
| 11 | self.use_multi_head = use_multi_head |
| 12 | self.n_heads = n_heads |
| 13 | |
| 14 | # Initialize intent prototypes |
| 15 | self.intent_prototypes = nn.Parameter( |
| 16 | nn.init.xavier_uniform_(torch.empty(n_intents, embed_dim)) |
| 17 | ) |
| 18 | |
| 19 | if use_multi_head: |
| 20 | self.head_dim = embed_dim // n_heads |
| 21 | self.attention_heads = nn.ModuleList([ |
| 22 | nn.Linear(embed_dim, self.head_dim) for _ in range(n_heads) |
| 23 | ]) |
| 24 | self.head_projections = nn.ModuleList([ |
| 25 | nn.Linear(embed_dim, self.head_dim) for _ in range(n_heads) |
| 26 | ]) |
| 27 | self.head_combine = nn.Linear(n_intents * n_heads, n_intents, bias=False) |
| 28 | |
| 29 | # Learnable temperature |
| 30 | self.temp_param = nn.Parameter(torch.ones(1) * temp) |
| 31 | |
| 32 | def compute_attention(self, embeddings): |
| 33 | if self.use_multi_head: |
nothing calls this directly
no outgoing calls
no test coverage detected