diff --git a/helm/modules/mice.py b/helm/modules/mice.py index c4fdaf1..02fdd3a 100644 --- a/helm/modules/mice.py +++ b/helm/modules/mice.py @@ -186,9 +186,9 @@ def __init__(self, manifold: Lorentz, args): self.curvature_list = np.linspace(0.1, 2.0, self.n_routed_experts).tolist() self.n_shared_experts = args.n_shared_experts # self.expert_manifolds = [Lorentz(c=(self.curvature_list[i]), learnable=False) for i in range(self.n_routed_experts)] - self.expert_manifolds = [Lorentz(c=1.0, learnable=args.train_curv) for i in range(self.n_routed_experts)] - self.experts = torch.nn.ModuleList([LorentzExpert(self.manifold, args.dim, args.moe_inter_dim, self.expert_manifolds[i]) for i in range(self.n_routed_experts)]) - self.shared_experts = LorentzFeedForward(self.manifold, args.dim, args.n_shared_experts * args.moe_inter_dim) + self.expert_manifolds = [Lorentz(c=1.0, learnable=bool(args.train_curv)) for i in range(self.n_routed_experts)] + self.experts = torch.nn.ModuleList([LorentzExpert(self.manifold, args.dim, args.mice_inter_dim, self.expert_manifolds[i]) for i in range(self.n_routed_experts)]) + self.shared_experts = LorentzFeedForward(self.manifold, args.dim, args.n_shared_experts * args.mice_inter_dim) if self.n_activated_experts == 2: self.add_experts = nn.LResNet(self.manifold, use_scale=True, scale=2.0, learn_scale=True) if self.n_shared_experts == 1: @@ -212,7 +212,9 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: """ shape = x.size() x = x.view(-1, self.dim) - weights, indices, scores = self.gate(x) + gate_out = self.gate(x) + weights, indices = gate_out[0], gate_out[1] + scores = gate_out[2] if len(gate_out) > 2 else None y = self.project(torch.zeros_like(x[..., 1:])) counts = torch.bincount(indices.flatten(), minlength=self.n_routed_experts).tolist() for i in range(self.experts_start_idx, self.experts_end_idx):