Skip to content

feat: add MOLE model support to DatasetSpecificMoEWrapper - #1848

Merged
misko merged 2 commits into
mainfrom
DatasetSpecificMoEWrapper_merge_head
Mar 4, 2026
Merged

misko merged 2 commits into
mainfrom
DatasetSpecificMoEWrapper_merge_head

Conversation

@misko

@misko misko commented Mar 4, 2026

Copy link
Copy Markdown
Contributor

Add merge_MOLE_model function to DatasetSpecificMoEWrapper for combining MOLE model heads with the main model during inference.

Add merge_MOLE_model function to DatasetSpecificMoEWrapper for combining
MOLE model heads with the main model during inference.
@meta-cla meta-cla Bot added the cla signed label Mar 4, 2026
@misko misko added enhancement New feature or request minor Minor version release labels Mar 4, 2026
@misko
misko marked this pull request as ready for review March 4, 2026 17:52
@misko
misko requested a review from rayg1234 March 4, 2026 17:52
self.global_mole_tensors.expert_mixing_coefficients = (
torch.zeros(1, len(self.dataset_name_to_exp), dtype=data.pos.dtype)
.scatter_(1, torch.tensor([[expert_idx]]), 1.0)
.to(data.pos.device)

@rayg1234 rayg1234 Mar 4, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

should the torch.zeros be initialized on device instead of moving it after?

nan_tensor = head_output[key].new_full(
head_output[key].shape, float("nan")
)
for dataset in self.non_merged_dataset_names:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

self.non_merged_dataset_names is None if not merged?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

correct, but guarded behind merged_on_dataset!=None
Its unclear in the correct form, i changed

  •    self.non_merged_dataset_names = None
    
  •    self.non_merged_dataset_names: list[str] = []
    

I think thats cleaner, does that work?

- Initialize expert_mixing_coefficients tensors directly on device instead
  of creating on CPU and transferring with .to(). Avoids unnecessary
  CPU→GPU transfer overhead.
- Change non_merged_dataset_names from None to empty list default.
  Makes iteration safe and type explicit with list[str] annotation.

@rayg1234 rayg1234 left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

lgtm!

@misko
misko added this pull request to the merge queue Mar 4, 2026
Merged via the queue into main with commit 9157056 Mar 4, 2026
15 checks passed
@misko
misko deleted the DatasetSpecificMoEWrapper_merge_head branch March 4, 2026 21:41
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

cla signed enhancement New feature or request minor Minor version release

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants