diff --git a/merging/main.py b/merging/main.py index 1bedcef..290e725 100644 --- a/merging/main.py +++ b/merging/main.py @@ -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() @@ -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: diff --git a/merging/taskloader.py b/merging/taskloader.py index c9170ba..b6464d4 100644 --- a/merging/taskloader.py +++ b/merging/taskloader.py @@ -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 '_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) @@ -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: @@ -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) @@ -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"] @@ -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) @@ -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: @@ -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) @@ -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"] @@ -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) @@ -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']