default k=0.5
This commit is contained in:
@@ -16,7 +16,7 @@ class NumberEmbedder(nn.Module):
|
|||||||
|
|
||||||
# 3) Comparator head: takes (ea, eb, e) -> logit for "a > b"
|
# 3) Comparator head: takes (ea, eb, e) -> logit for "a > b"
|
||||||
class PairwiseComparator(nn.Module):
|
class PairwiseComparator(nn.Module):
|
||||||
def __init__(self, d=4, hidden=16, k=1.0):
|
def __init__(self, d=4, hidden=16, k=0.5):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.log_k = nn.Parameter(torch.tensor([k]))
|
self.log_k = nn.Parameter(torch.tensor([k]))
|
||||||
self.embed = NumberEmbedder(d, hidden)
|
self.embed = NumberEmbedder(d, hidden)
|
||||||
|
|||||||
Reference in New Issue
Block a user