Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 3 additions & 6 deletions merging/main.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,7 @@
from prepare_args import prepare_args, create_parser
from taskloader import get_ft_ckpts
import importlib

def get_ft_ckpts(base_model):
model_name = base_model.split('/')[-1]
task_names = ['instruction', 'math', 'coding', 'safety', 'multilingual']
return [f'MergeBench/{model_name}_{task_name}' for task_name in task_names]

def parse_args():
parser = create_parser()

Expand All @@ -18,7 +14,8 @@ def parse_args():
def main(args):
kwargs = prepare_args(args)
merger_module = importlib.import_module("merging_methods")
ft_ckpts = get_ft_ckpts(args.base_model)
task_names = args.task_names.split('-') if args.task_names else None
ft_ckpts = get_ft_ckpts(args.base_model, task_names)

kwargs_str = "_".join(f"{key}_{value}" for key, value in kwargs.items() if key not in ['fisher_only','merge_only','save_group','task_names','keep_checkpoints'])
if args.save_group:
Expand Down
56 changes: 48 additions & 8 deletions merging/taskloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,18 +42,50 @@ def formatting_prompts_func(examples, instruction_key='instruction', input_key='


class TaskLoader:
# Category name shared by a task's fine-tuned checkpoint suffix and its
# validation set (e.g. 'math' -> checkpoint '<model>_math', data 'MergeBench/math_val').
# Overridden by each concrete task; this is the single source of truth that keeps
# the merged checkpoint, its Gram statistics, and the task ordering aligned.
category = None

def __new__(cls, task_name, *args, **kwargs):
if task_name in globals() and issubclass(globals()[task_name], cls):
subclass = globals()[task_name]
return super().__new__(subclass)
subclass = globals()[task_name]
return super().__new__(subclass)
else:
raise ValueError(f"Invalid task name: {task_name}")

def __init__(self, task_name, *args, **kwargs):
self.task_name = task_name

@classmethod
def checkpoint_for(cls, model_name):
return f'MergeBench/{model_name}_{cls.category}'


# Canonical task ordering, used as a fallback for merging algorithms that do not
# take an explicit --task_names argument (their operation is symmetric across models,
# so the order does not affect correctness).
DEFAULT_TASK_ORDER = ['Tulu3IF', 'DartMath', 'MagiCoder', 'WildguardMix', 'Aya']


def get_ft_ckpts(base_model, task_names=None):
"""Derive the fine-tuned checkpoint paths from the task ordering.

task_names is the ordered list of TaskLoader class names (e.g. from
--task_names.split('-')). Deriving the checkpoints from the same list the
merge loop iterates guarantees each checkpoint lines up with the dataset used
to compute its Gram statistics. Falls back to DEFAULT_TASK_ORDER when no
task_names are given.
"""
model_name = base_model.split('/')[-1]
task_names = task_names or DEFAULT_TASK_ORDER
return [globals()[task_name].checkpoint_for(model_name) for task_name in task_names]


class WildguardMix(TaskLoader):
category = 'safety'

def __init__(self, task_name, model, tokenizer, sample_size=None):
super().__init__(task_name, model, tokenizer, sample_size=sample_size)

Expand All @@ -73,7 +105,7 @@ def __init__(self, task_name, model, tokenizer, sample_size=None):
save_strategy='no',
)

self.training_dataset = load_dataset('MergeBench/safety_val',cache_dir=cache_dir)
self.training_dataset = load_dataset(f'MergeBench/{self.category}_val',cache_dir=cache_dir)
self.training_dataset = self.training_dataset.rename_column("prompt", "query")

if sample_size is None:
Expand All @@ -90,6 +122,8 @@ def __init__(self, task_name, model, tokenizer, sample_size=None):


class MagiCoder(TaskLoader):
category = 'coding'

def __init__(self, task_name, model, tokenizer, sample_size=None):
super().__init__(task_name, model, tokenizer, sample_size=sample_size)

Expand All @@ -109,7 +143,7 @@ def __init__(self, task_name, model, tokenizer, sample_size=None):
save_strategy='no',
)

self.training_dataset = load_dataset('MergeBench/coding_val',cache_dir=cache_dir)
self.training_dataset = load_dataset(f'MergeBench/{self.category}_val',cache_dir=cache_dir)

if sample_size is None:
self.training_dataset = self.training_dataset["train"]
Expand All @@ -125,6 +159,8 @@ def __init__(self, task_name, model, tokenizer, sample_size=None):


class Aya(TaskLoader):
category = 'multilingual'

# TODO: match with Yuzheng's config
def __init__(self, task_name, model, tokenizer, sample_size=None):
super().__init__(task_name, model, tokenizer, sample_size=sample_size)
Expand All @@ -145,7 +181,7 @@ def __init__(self, task_name, model, tokenizer, sample_size=None):
save_strategy='no',
)

self.training_dataset = load_dataset('MergeBench/multilingual_val',cache_dir=cache_dir)
self.training_dataset = load_dataset(f'MergeBench/{self.category}_val',cache_dir=cache_dir)
if sample_size is None:
self.training_dataset = self.training_dataset["train"]
else:
Expand All @@ -160,6 +196,8 @@ def __init__(self, task_name, model, tokenizer, sample_size=None):


class DartMath(TaskLoader):
category = 'math'

def __init__(self, task_name, model, tokenizer, sample_size=None):
super().__init__(task_name, model, tokenizer, sample_size=sample_size)

Expand All @@ -179,7 +217,7 @@ def __init__(self, task_name, model, tokenizer, sample_size=None):
save_strategy='no',
)

self.training_dataset = load_dataset('MergeBench/math_val',cache_dir=cache_dir)
self.training_dataset = load_dataset(f'MergeBench/{self.category}_val',cache_dir=cache_dir)

if sample_size is None:
self.training_dataset = self.training_dataset["train"]
Expand All @@ -194,6 +232,8 @@ def __init__(self, task_name, model, tokenizer, sample_size=None):
)

class Tulu3IF(TaskLoader):
category = 'instruction'

def __init__(self, task_name, model, tokenizer, sample_size=None):
super().__init__(task_name, model, tokenizer, sample_size=sample_size)

Expand All @@ -213,7 +253,7 @@ def __init__(self, task_name, model, tokenizer, sample_size=None):
save_strategy='no',
)

self.training_dataset = load_dataset('MergeBench/instruction_val',cache_dir=cache_dir)
self.training_dataset = load_dataset(f'MergeBench/{self.category}_val',cache_dir=cache_dir)

if sample_size is None:
self.training_dataset = self.training_dataset['train']
Expand Down