diff --git a/src/xlm/commands/lightning_train.py b/src/xlm/commands/lightning_train.py index 2bfcabbf..f5027bf1 100644 --- a/src/xlm/commands/lightning_train.py +++ b/src/xlm/commands/lightning_train.py @@ -18,6 +18,10 @@ from lightning.pytorch.loggers import Logger from lightning import Callback from transformers.modeling_utils import no_init_weights +from xlm.utils.model_loading import ( + load_model_weights_into_model, + _get_model_only_checkpoint_path, +) from xlm.utils.rank_zero import RankedLogger logger = RankedLogger(__name__, rank_zero_only=True) @@ -111,21 +115,16 @@ def train(cfg: DictConfig): ckpt_path = os.path.join(cfg.checkpointing_dir, "last.ckpt") if not os.path.isfile(ckpt_path): ckpt_path = None - # check if we have model only checkpoint - model_only_ckpt_path = None - if cfg.get("model_only_checkpoint_path", None) is not None: - if ckpt_path is not None: + + if ckpt_path is not None: + has_model_only = cfg.get("model_only_checkpoint_path", None) is not None + has_hub = OmegaConf.select(cfg, "hub.repo_id", default=None) is not None + if has_model_only or has_hub: logger.error( - "model_only_checkpoint_path and resume_from_checkpoint cannot both be provided." - " We will use the resume_from_checkpoint path " - f"{ckpt_path} for the model weights as well." + "Resume checkpoint is set; model-only / Hub weight sources are ignored. " + f"Using {ckpt_path} for model weights." ) - else: - if not os.path.isfile(cfg.model_only_checkpoint_path): - raise ValueError( - f"The model only checkpoint path {cfg.model_only_checkpoint_path} does not exist." - ) - model_only_ckpt_path = cfg.model_only_checkpoint_path + model_only_ckpt_path = _get_model_only_checkpoint_path(cfg, "", True, ckpt_path) logger.info(f"Instantiating trainer <{cfg.trainer._target_}>") trainer = hydra.utils.instantiate( @@ -169,14 +168,17 @@ def train(cfg: DictConfig): _recursive_=False, ) if model_only_ckpt_path is not None: - message = lightning_module.model.load_state_dict( - torch.load(model_only_ckpt_path) - ) logger.warning( - "Loading weights for `model` from a pretrained model at " - f"{model_only_ckpt_path} before call to `trainer.fit` => before `.setup()` and `.configure_model()`" + "Loading weights for `model` from " + f"{model_only_ckpt_path} before `trainer.fit` (before `.setup()` / `.configure_model()`)." + ) + load_model_weights_into_model( + lightning_module.model, + model_only_ckpt_path, + map_location="cpu", + strict=True, + weights_only=True, ) - logger.warning(message) # train if cfg.job_type == "train": diff --git a/src/xlm/configs/lightning_train/datasets/humaneval_infill_test.yaml b/src/xlm/configs/lightning_train/datasets/humaneval_infill_test.yaml new file mode 100644 index 00000000..85dffeba --- /dev/null +++ b/src/xlm/configs/lightning_train/datasets/humaneval_infill_test.yaml @@ -0,0 +1,14 @@ +defaults: + - default_eval + +collator: ??? +load_func: xlm.tasks.humaneval_infill_task.load_func +load_func_kwargs: + benchmark_name: ??? +preprocess_function: xlm.tasks.humaneval_infill_task.humaneval_infill_preprocess_fn +full_name: null +columns_to_keep: + - canonical_solution + - prompt + - suffix + - task_id \ No newline at end of file diff --git a/src/xlm/configs/lightning_train/datasets/opencoder_train.yaml b/src/xlm/configs/lightning_train/datasets/opencoder_train.yaml new file mode 100644 index 00000000..3fb80431 --- /dev/null +++ b/src/xlm/configs/lightning_train/datasets/opencoder_train.yaml @@ -0,0 +1,24 @@ +_target_: xlm.datamodule.LocalDatasetManager +ds_type: "parquet" +collator: ??? # model specific collator +full_name: opencoder/train +full_name_debug: opencoder/train +load_kwargs: + data_files: ??? +preprocess_function: xlm.tasks.opencoder.opencoder_preprocess_fn +preprocess_function_kwargs: + prompt_key: prompt + response_key: response +on_the_fly_processor: xlm.datamodule.token_ids_to_input_ids +on_the_fly_group_processor: null +stages: + - fit +dataloader_kwargs: + batch_size: ${per_device_batch_size} # per device, depends on the device type + num_workers: ${num_dataloader_workers} + shuffle: null # can't specify shuffle for IterableDataset + pin_memory: True + persistent_workers: False # INFO: see: https://github.com/huggingface/datasets/issues/7447 + prefetch_factor: ${dataloader_prefetch_factor} + drop_last: True +model_name: null \ No newline at end of file diff --git a/src/xlm/configs/lightning_train/datasets/opencoder_val.yaml b/src/xlm/configs/lightning_train/datasets/opencoder_val.yaml new file mode 100644 index 00000000..55f6a808 --- /dev/null +++ b/src/xlm/configs/lightning_train/datasets/opencoder_val.yaml @@ -0,0 +1,25 @@ +_target_: xlm.datamodule.LocalDatasetManager +ds_type: "parquet" +collator: ??? # model specific collator +full_name: opencoder/val +full_name_debug: opencoder/val +load_kwargs: + data_files: ??? +preprocess_function: xlm.tasks.opencoder.opencoder_preprocess_fn +preprocess_function_kwargs: + prompt_key: prompt + response_key: response +on_the_fly_processor: xlm.datamodule.token_ids_to_input_ids +on_the_fly_group_processor: null +stages: + - fit + - validate +dataloader_kwargs: + batch_size: ${per_device_batch_size} # per device, depends on the device type + num_workers: ${num_dataloader_workers} + shuffle: null # can't specify shuffle for IterableDataset + pin_memory: True + persistent_workers: False # INFO: see: https://github.com/huggingface/datasets/issues/7447 + prefetch_factor: ${dataloader_prefetch_factor} + drop_last: True +model_name: null \ No newline at end of file diff --git a/src/xlm/datamodule.py b/src/xlm/datamodule.py index 704e095c..f4c6ba5f 100644 --- a/src/xlm/datamodule.py +++ b/src/xlm/datamodule.py @@ -609,7 +609,15 @@ def from_txt(cls, txt_file: Union[str, os.PathLike], **kwargs): vocab.append(line.strip()) return cls(vocab=vocab, **kwargs) - +class SimpleSpaceTokenizerWithDeletion(SimpleSpaceTokenizer): + + def __init__(self, vocab: Sequence[str], **kwargs): + super().__init__(vocab=vocab, **kwargs) + del_token_id = len(self._vocab_str_to_int) + self._vocab_str_to_int["[DEL]"] = del_token_id + self._vocab_int_to_str[del_token_id] = "[DEL]" + setattr(self, "delete_token", "[DEL]") + setattr(self, "delete_token_id", del_token_id) class SimpleSpaceTokenizerWithCyclicPads(SimpleSpaceTokenizer): """SimpleSpaceTokenizer with cyclic pad tokens (pad_0..pad_{n-1}).""" @@ -1280,6 +1288,21 @@ def _download(self, num_proc: Optional[int] = None) -> datasets.Dataset: num_proc=num_proc, )["train"] return ds + elif self.ds_type == "parquet": + load_kwargs_copy = self.load_kwargs.copy() + if "data_files" in load_kwargs_copy: + data_files = load_kwargs_copy.pop("data_files") + else: + file_name = f"{self._split_to_download}.parquet" + _path = Path(self.full_name).parent + data_files = str(_path / file_name) + ds = datasets.load_dataset( + "parquet", + data_files=data_files, + **load_kwargs_copy, + num_proc=num_proc, + )['train'] + return ds else: raise ValueError(f"Unsupported dataset type: {self.ds_type}") @@ -1313,9 +1336,13 @@ def __init__( stages: Optional[ List[Literal["fit", "validate", "test", "predict"]] ] = None, + load_func: Optional[str] = None, + load_func_kwargs: Optional[Dict[str, Any]] = None, ): self.collator = collator self.full_name = full_name + self.load_func = load_func + self.load_func_kwargs = load_func_kwargs or {} self.dataloader_kwargs = dataloader_kwargs self.preprocess_function = preprocess_function self.preprocess_function_kwargs = preprocess_function_kwargs or {} @@ -1404,7 +1431,13 @@ def prepare_data( logger.info( f"EvalDatasetManager: preparing {self.full_name} (no manual cache)" ) - ds = self._download(num_proc=num_proc) + if self.load_func: + load_fn: Callable[..., Any] = get_function( + self.load_func + ) + ds = load_fn(**self.load_func_kwargs) + else: + ds = self._download(num_proc=num_proc) ds = self._preprocess(ds, tokenizer, num_proc=num_proc) return ds diff --git a/src/xlm/external_models.py b/src/xlm/external_models.py index 66e1b0d8..0d141713 100644 --- a/src/xlm/external_models.py +++ b/src/xlm/external_models.py @@ -27,7 +27,7 @@ ENV_XLM_MODELS_PATH = "XLM_MODELS_PATH" # dir containing external models ENV_XLM_MODELS_PACKAGES = "XLM_MODELS_PACKAGES" # installed python packages containing external models, comma separated list of package names CORE_XLM_MODELS = ( - "arlm:mlm:ilm:mdlm:flexmdm" # core models available in xlm-models package + "arlm:mlm:ilm:mdlm:flexmdm:dream:dreamon" # core models available in xlm-models package ) diff --git a/src/xlm/tasks/humaneval_infill_task.py b/src/xlm/tasks/humaneval_infill_task.py new file mode 100644 index 00000000..5b41ef78 --- /dev/null +++ b/src/xlm/tasks/humaneval_infill_task.py @@ -0,0 +1,245 @@ +"""HumanEval-Infill (single-line) for xlm eval / DreamOn-style runs. +""" + +from __future__ import annotations + +import json +import os +import tempfile +import abc +from typing import Any, Dict, List, Optional, Tuple +from transformers import AutoTokenizer +from xlm.utils.rank_zero import RankedLogger +from human_eval_infilling.data import read_problems +from human_eval_infilling.evaluation import evaluate_functional_correctness +from datasets import Dataset + +logger = RankedLogger(__name__, rank_zero_only=True) + + +def humaneval_infill_preprocess_fn( + example: Dict[str, Any], + tokenizer: Any, +) -> Dict[str, Any]: + """Tokenize prefix / suffix / middle and build a single span of mask tokens. + + Returns ``prompt_ids`` (prefix + masks + suffix) and ``input_ids`` (full + sequence without masks) for :class:`mlm.datamodule_mlm.MLMInfillWithExactTargetPredCollator`. + """ + prefix = example["prompt"] + suffix = example["suffix"] + middle = example["canonical_solution"] + task_id = example["task_id"] + + pre_ids = tokenizer.encode(prefix) + suf_ids = tokenizer.encode(suffix) + mid_ids = tokenizer.encode(middle) + + return { + "prefix_ids": pre_ids, + "suffix_ids": suf_ids, + "middle_ids": mid_ids, + "task_id": task_id, + "canonical_solution": middle, + "prefix": prefix, + "suffix": suffix, + } + +class Tokenizer(abc.ABC): + @abc.abstractmethod + def encode(self, tokens, add_bos, add_eos): + pass + + @abc.abstractmethod + def decode(self, tokens): + pass + + @abc.abstractmethod + def get_token_offsets( + self, text: str, tokens: Optional[List[int]] = None + ) -> Tuple[List[str], List[int]]: + """Return the offsets of the tokens in the original text. Only used for evaluation.""" + pass + +class HFTokenizerWrapper(Tokenizer): + def __init__(self, hf_tokenizer: str) -> None: + self.tokenizer = hf_tokenizer + self.bos_id = self.tokenizer.bos_token_id + self.eos_id = self.tokenizer.eos_token_id + self.mask_id = self.tokenizer.mask_token_id + self.pad_id = self.tokenizer.pad_token_id + self.expand_id = 151667 + + self.bos_token_id = self.bos_id + self.eos_token_id = self.eos_id + self.mask_token_id = self.mask_id + self.expand_token_id = self.expand_id + self.pad_token_id = self.pad_id + + def encode(self, s: str, add_bos: bool = False, add_eos: bool = False): + tokens = [self.bos_id] * add_bos + self.tokenizer.encode(s) + [self.eos_id] * add_eos + return tokens + + def decode(self, tokens: List[int], **kwargs): + return self.tokenizer.decode(tokens, **kwargs) + + def get_token_offsets( + self, text: str, tokens: Optional[List[int]] = None + ) -> Tuple[List[str], List[int]]: + """Return the offsets of the tokens in the original text. Only used for evaluation.""" + pass + +def get_tokenizer(pretrained_model_name_or_path: str ): + tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path, trust_remote_code=True) + tokenizer = HFTokenizerWrapper(tokenizer) + return tokenizer + +def load_func(benchmark_name: str): + problems = read_problems(benchmark_name) + prefixs = [problems[task_id]["prompt"] for task_id in problems] + suffixs = [problems[task_id]["suffix"] for task_id in problems] + canonical_solutions = [problems[task_id]["canonical_solution"] for task_id in problems] + test = [problems[task_id]["test"] for task_id in problems] + entry_points = [problems[task_id]["entry_point"] for task_id in problems] + data = { + "task_id": list(problems.keys()), + "prompt": prefixs, + "suffix": suffixs, + "canonical_solution": canonical_solutions, + "test": test, + "entry_points": entry_points, + } + return Dataset.from_dict(data) + +class HumanEvalInfillEval: + """Post-hoc evaluator: write HumanEval-Infill samples and optional pass@k. + + ``predictions`` entries should include ``text`` (full decoded infill line), + ``prefix``, ``suffix``, and ``task_id`` (from ``additional_fields_from_batch``). + """ + + def __init__(self,benchmark_name): + self.benchmark_name = benchmark_name + + def eval( + self, + predictions: List[Dict[str, Any]], + tokenizer: Any = None, + **kwargs: Any, + ) -> Tuple[List[Dict[str, Any]], Dict[str, Any]]: + del tokenizer, kwargs + if not predictions: + return predictions, {} + + samples: List[Dict[str, Any]] = [] + for pred in predictions: + full = pred.get("text", "") or "" + prefix = pred.get("prefix", "") or "" + suffix = pred.get("suffix", "") or "" + task_id = pred.get("task_id", "") + samples.append( + { + "task_id": task_id, + "completion": full, + "prefix": prefix, + "suffix": suffix, + "ground_truth_middle": pred.get("canonical_solution", ""), + } + ) + + metrics: Dict[str, Any] = {} + + with tempfile.NamedTemporaryFile( + mode="w", suffix=".jsonl", delete=False, encoding="utf-8" + ) as f: + for row in samples: + f.write(json.dumps(row) + "\n") + tmp_path = f.name + try: + results = evaluate_functional_correctness( + self.benchmark_name, + tmp_path, + [1], + n_workers=int(os.environ.get("HUMANEVAL_WORKERS", "16")), + timeout=float(os.environ.get("HUMANEVAL_TIMEOUT", "3.0")), + ) + metrics = results + logger.info("HumanEvalInfillEval: %s", results) + except Exception as e: + print( + f"HumanEvalInfillEval: evaluate_functional_correctness failed: {e}" + ) + logger.exception( + "HumanEvalInfillEval: evaluate_functional_correctness failed" + ) + finally: + try: + os.unlink(tmp_path) + except OSError: + pass + + return predictions, metrics + + +def _jsonl_row_to_preprocess_example(row: Dict[str, Any]) -> Dict[str, Any]: + """Map a JSONL object (benchmark or prediction row) to ``humaneval_infill_preprocess_fn`` inputs.""" + prompt = row.get("prompt") + if prompt is None: + prompt = row.get("prefix", "") or "" + canon = ( + row.get("canonical_solution") + or row.get("ground_truth_middle") + or row.get("middle") + or "" + ) + return { + "task_id": row.get("task_id", ""), + "prompt": prompt, + "suffix": row.get("suffix", "") or "", + "ground_truth_middle": canon, + "completion": row.get("text", ""), + } + + +def humaneval_infill_from_file(file_path: str,benchmark_name: str) -> List[Dict[str, Any]]: + import gzip + + rows_out: List[Dict[str, Any]] = [] + if file_path.endswith(".gz"): + fp_ctx = gzip.open(file_path, "rt", encoding="utf-8") + else: + fp_ctx = open(file_path, "r", encoding="utf-8") + with fp_ctx as fp: + for line in fp: + raw = json.loads(line) + example = _jsonl_row_to_preprocess_example(raw) + rows_out.append(example) + with tempfile.NamedTemporaryFile( + mode="w", suffix=".jsonl", delete=False, encoding="utf-8" + ) as f: + for row in rows_out: + f.write(json.dumps(row) + "\n") + tmp_path = f.name + try: + results = evaluate_functional_correctness( + benchmark_name, + tmp_path, + [1], + n_workers=int(os.environ.get("HUMANEVAL_WORKERS", "16")), + timeout=float(os.environ.get("HUMANEVAL_TIMEOUT", "3.0")), + ) + metrics = results + logger.info("HumanEvalInfillEval: %s", results) + except Exception as e: + print( + f"HumanEvalInfillEval: evaluate_functional_correctness failed: {e}" + ) + logger.exception( + "HumanEvalInfillEval: evaluate_functional_correctness failed" + ) + finally: + try: + os.unlink(tmp_path) + except OSError: + pass + return metrics \ No newline at end of file diff --git a/src/xlm/tasks/opencoder.py b/src/xlm/tasks/opencoder.py new file mode 100644 index 00000000..b6d855ae --- /dev/null +++ b/src/xlm/tasks/opencoder.py @@ -0,0 +1,63 @@ +from typing import Any, Dict +from transformers import AutoTokenizer +import re +from xlm.tasks.humaneval_infill_task import HFTokenizerWrapper + +def extract_code_block(response: str): + """ + Extracts the content inside a markdown-style Python code block: + + Returns: + prefix: everything before ```python + code_block: content inside the code block + suffix: everything after ``` + """ + # Split by ```python and ``` to extract the code block + parts = re.split(r'```python\s*|\s*```', response) + + if len(parts) < 2: + return "", response, "" # No code block found + + prefix = parts[0] + code_block = '```' + parts[1] + '\n```' if len(parts) > 1 else "" + suffix = parts[2] if len(parts) > 2 else "" + + return prefix, code_block, suffix + +def opencoder_preprocess_fn( + example: Dict[str, Any], + tokenizer: Any, + prompt_key: str, + response_key: str, +) -> Dict[str, Any]: + """Tokenize prompt and response.""" + prompt = example[prompt_key] + response = example[response_key] + if not isinstance(prompt, str): + prompt_chat = list(prompt) + else: + prompt_chat = [{"role": "user", "content": prompt}] + + # string + prompt_chat_str = tokenizer.apply_chat_template( + prompt_chat, add_generation_prompt=True, tokenize=False + ) + response_chat_str = response + tokenizer.eos_token + + prompt_ids_output = tokenizer.encode(prompt_chat_str) + response_ids_output = tokenizer.encode(response_chat_str) + prefix, code_block, suffix = extract_code_block(response_chat_str) + + return { + "prompt_ids": prompt_ids_output, + "token_ids": response_ids_output, + "prefix": prefix, + "middle": code_block, + "suffix": suffix, + } + + +def get_tokenizer(pretrained_model_name_or_path: str): + tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path,trust_remote_code=True) + tokenizer = HFTokenizerWrapper(tokenizer) + return tokenizer \ No newline at end of file diff --git a/src/xlm/tasks/opencoder_prep_data.py b/src/xlm/tasks/opencoder_prep_data.py new file mode 100644 index 00000000..f82dcc68 --- /dev/null +++ b/src/xlm/tasks/opencoder_prep_data.py @@ -0,0 +1,72 @@ +from datasets import load_dataset, Dataset +from concurrent.futures import ProcessPoolExecutor +import tqdm +import os +import pandas as pd + +# Global tokenizer (must be initialized in child processes) +tokenizer = None + +def init_tokenizer(model_path): + global tokenizer + from transformers import AutoTokenizer + tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code = True) + +def tokenize_and_filter_row(row): + try: + prompt_len = len(tokenizer.tokenize(row['prompt'])) + response_len = len(tokenizer.tokenize(row['response'])) + return (prompt_len + response_len) <= 1024 + except Exception as e: + print(f"Error tokenizing row: {e}") + return False + +def main(): + # Step 1: Load and rename columns + output_dir = '/scratch4/workspace/dmaram_umass_edu-dmaram/learned-correctors/data/opencoder-stage2-edu' + data = load_dataset("OpenCoder-LLM/opc-sft-stage2", "educational_instruct") + data = data.rename_column("instruction", "prompt") + data = data.rename_column("output", "response") + + train_data = data["train"] + + model_path = "Dream-org/Dream-Coder-v0-Base-7B" + + # Step 2: Filter using multiprocessing + records = [dict(train_data[i]) for i in range(len(train_data))] + + with ProcessPoolExecutor(initializer=init_tokenizer, initargs=(model_path,)) as executor: + results = list(tqdm.tqdm(executor.map(tokenize_and_filter_row, records), total=len(records))) + + filtered_indices = [i for i, keep in enumerate(results) if keep] + print(f"Total filtered samples (<=1024 tokens): {len(filtered_indices)}") + + filtered_dataset = train_data.select(filtered_indices) # This is still a Dataset + + # Step 3: Shuffle and split + filtered_dataset = filtered_dataset.shuffle(seed=42) + + eval_size = 1000 + total_size = len(filtered_dataset) + if total_size < eval_size: + print(f"Warning: Only {total_size} samples available. Using all for training.") + train_dataset = filtered_dataset + eval_dataset = Dataset.from_dict({}) # empty dataset + else: + eval_dataset = filtered_dataset.select(range(eval_size)) + train_dataset = filtered_dataset.select(range(eval_size, total_size)) + + # Step 4: Convert to Pandas and save + os.makedirs(output_dir, exist_ok=True) + + # Convert to Pandas before saving + train_df = train_dataset.to_pandas() + eval_df = eval_dataset.to_pandas() + + train_df.to_parquet(os.path.join(output_dir, "train_data.parquet"), index=False) + eval_df.to_parquet(os.path.join(output_dir, "eval_data.parquet"), index=False) + + print(f"Saved {len(train_df)} train samples and {len(eval_df)} eval samples.") + +if __name__ == '__main__': + main() \ No newline at end of file diff --git a/src/xlm/utils/model_loading.py b/src/xlm/utils/model_loading.py index d5d5f9e4..cf579a54 100644 --- a/src/xlm/utils/model_loading.py +++ b/src/xlm/utils/model_loading.py @@ -1,7 +1,8 @@ -"""Unified model loading for inference across commands. +"""Unified model loading for inference and optional train-time weight bootstrap. This module provides a single, consistent interface for loading models -for inference tasks (generation, evaluation, push to hub, demos). +for inference tasks (generation, evaluation, push to hub, demos) and for +resolving Hub/local model-only checkpoints used by ``lightning_train``. """ import contextlib diff --git a/xlm-models/dreamon/configs/collator/humaneval_infill_pred_dreamon.yaml b/xlm-models/dreamon/configs/collator/humaneval_infill_pred_dreamon.yaml new file mode 100644 index 00000000..75ce880a --- /dev/null +++ b/xlm-models/dreamon/configs/collator/humaneval_infill_pred_dreamon.yaml @@ -0,0 +1,15 @@ +_target_: mlm.datamodule_mlm.MLMDreamInfillPredCollator +tokenizer: ${global_components:tokenizer} +add_bos: false +add_eos: false +mask_expansion: true +max_tokens: 2048 +max_prompt_len: 2048 +min_gen_len: 64 +max_gen_len: 64 +pad_to_max_len: false +pass_through_fields: + - canonical_solution + - prefix + - suffix + - task_id diff --git a/xlm-models/dreamon/configs/collator/infill_dreamon.yaml b/xlm-models/dreamon/configs/collator/infill_dreamon.yaml new file mode 100644 index 00000000..aff251e2 --- /dev/null +++ b/xlm-models/dreamon/configs/collator/infill_dreamon.yaml @@ -0,0 +1,12 @@ +_target_: dreamon.datamodule_dreamon.DreamOnInfillTrainCollator +tokenizer: ${global_components:tokenizer} +max_length: 1024 +middle_strategy: line +middle_line_num: null +truncation: right +merge_prob: 0.5 +max_delete: 64 +merge_schedule: dynamic_inverse +use_uniform_merge_prob: 0.5 +expand_token_id: 151667 + diff --git a/xlm-models/dreamon/configs/datamodule/humaneval_infill_dreamon.yaml b/xlm-models/dreamon/configs/datamodule/humaneval_infill_dreamon.yaml new file mode 100644 index 00000000..b22a213b --- /dev/null +++ b/xlm-models/dreamon/configs/datamodule/humaneval_infill_dreamon.yaml @@ -0,0 +1,22 @@ +# @package _global_ +defaults: + - default + - /collator@datamodule.dataset_managers.val.humaneval_infill_prediction.collator: humaneval_infill_pred_dreamon + - /collator@datamodule.dataset_managers.test.humaneval_infill_prediction.collator: humaneval_infill_pred_dreamon + - /datasets@datamodule.dataset_managers.val.humaneval_infill_prediction: humaneval_infill_test + - /datasets@datamodule.dataset_managers.test.humaneval_infill_prediction: humaneval_infill_test + +datamodule: + print_batch_fn: dream.datamodule_dream.print_batch_dream + dataset_managers: + test: + humaneval_infill_prediction: + load_func_kwargs: + benchmark_name: test + val: + humaneval_infill_prediction: + load_func_kwargs: + benchmark_name: test + +tags: + dataset: humaneval_infill diff --git a/xlm-models/dreamon/configs/datamodule/opencoder_dreamon.yaml b/xlm-models/dreamon/configs/datamodule/opencoder_dreamon.yaml new file mode 100644 index 00000000..72930ba9 --- /dev/null +++ b/xlm-models/dreamon/configs/datamodule/opencoder_dreamon.yaml @@ -0,0 +1,21 @@ +# @package _global_ +defaults: + - default + - /collator@datamodule.dataset_managers.train.lm.collator: infill_dreamon + - /collator@datamodule.dataset_managers.val.lm.collator: infill_dreamon + - /datasets@datamodule.dataset_managers.train.lm: opencoder_train + - /datasets@datamodule.dataset_managers.val.lm: opencoder_val + - /collator@datamodule.dataset_managers.test.humaneval_infill_prediction.collator: humaneval_infill_pred_dreamon + - /datasets@datamodule.dataset_managers.test.humaneval_infill_prediction: humaneval_infill_test + +datamodule: + print_batch_fn: dream.datamodule_dream.print_batch_dream + dataset_managers: + test: + humaneval_infill_prediction: + load_func_kwargs: + benchmark_name: test + +tags: + dataset: opencoder_dreamon + diff --git a/xlm-models/dreamon/configs/experiment/humaneval_infill_dreamon_eval.yaml b/xlm-models/dreamon/configs/experiment/humaneval_infill_dreamon_eval.yaml new file mode 100644 index 00000000..2bb01f81 --- /dev/null +++ b/xlm-models/dreamon/configs/experiment/humaneval_infill_dreamon_eval.yaml @@ -0,0 +1,60 @@ +# @package _global_ +defaults: + - override /callbacks: debug_v2 + - override /datamodule: humaneval_infill_dreamon + - override /noise_schedule: dummy + - override /model_type: dreamon_eval + - override /model: dreamon_7b + +per_device_batch_size: 1 +global_batch_size: 1 +monitored_metric: null +init_dtype: bfloat16 +skip_init_weights: true + +hub: + repo_id: Dream-org/DreamOn-v0-7B + +eval: + split: test + +global_components: + tokenizer: + _target_: xlm.tasks.humaneval_infill_task.get_tokenizer + pretrained_model_name_or_path: Dream-org/DreamOn-v0-7B + +predictor: + mask_expansion: true + diffusion_kwargs: + dtype: bf16 + steps: 256 + max_gen_len: 64 + min_gen_len: 64 + alg: entropy + batch_size: 1 + pad_to_max_len: false + alg_temp: 0.0 + temperature: 0.2 + top_p: 0.9 + delete_eos_token: true + pad_eos_to_right: true + show_progress: false + max_tokens: 2048 + +log_predictions: + _target_: xlm.log_predictions.LogPredictions + writers: + - file + additional_fields_from_batch: + - task_id + - prefix + - suffix + - canonical_solution + +post_hoc_evaluator: + _target_: xlm.tasks.humaneval_infill_task.HumanEvalInfillEval + benchmark_name: single-line + +trainer: + max_steps: 0 + num_sanity_val_steps: 0 diff --git a/xlm-models/dreamon/configs/experiment/opencoder_dreamon.yaml b/xlm-models/dreamon/configs/experiment/opencoder_dreamon.yaml new file mode 100644 index 00000000..e6e731ae --- /dev/null +++ b/xlm-models/dreamon/configs/experiment/opencoder_dreamon.yaml @@ -0,0 +1,60 @@ +# @package _global_ +defaults: + - override /callbacks: debug_v2 + - override /datamodule: opencoder_dreamon + - override /noise_schedule: dummy + - override /model_type: dreamon + - override /model: dreamon_7b # Backbone is the same for Dream and DreamCoder + +per_device_batch_size: 1 +global_batch_size: 2 +monitored_metric: null +init_dtype: bfloat16 +skip_init_weights: true + +hub: + repo_id: Dream-org/Dream-Coder-v0-Base-7B + +global_components: + tokenizer: + _target_: xlm.tasks.opencoder.get_tokenizer + pretrained_model_name_or_path: Dream-org/Dream-Coder-v0-Base-7B + +predictor: + mask_expansion: true + diffusion_kwargs: + dtype: bf16 + steps: 256 + max_gen_len: 64 + min_gen_len: 64 + alg: entropy + batch_size: 1 + pad_to_max_len: false + alg_temp: 0.0 + temperature: 0.2 + top_p: 0.9 + delete_eos_token: true + pad_eos_to_right: true + show_progress: false + max_tokens: 2048 + +log_predictions: + _target_: xlm.log_predictions.LogPredictions + writers: + - file + additional_fields_from_batch: + - task_id + - prefix + - suffix + - canonical_solution + +post_hoc_evaluator: + _target_: xlm.tasks.humaneval_infill_task.HumanEvalInfillEval + benchmark_name: test + +trainer: + max_steps: 2 + val_check_interval: 1 + num_sanity_val_steps: 0 + check_val_every_n_epoch: 1 + limit_val_batches: 1 diff --git a/xlm-models/dreamon/configs/model/dreamon_7b.yaml b/xlm-models/dreamon/configs/model/dreamon_7b.yaml new file mode 100644 index 00000000..c3a1f722 --- /dev/null +++ b/xlm-models/dreamon/configs/model/dreamon_7b.yaml @@ -0,0 +1,26 @@ +_target_: dreamon.dreamon_model.DreamOnModel +config: + _target_: dreamon.configuration_dreamon.DreamOnConfig + expand_token_id: 151667 + vocab_size: 152064 + hidden_size: 3584 + intermediate_size: 18944 + num_hidden_layers: 28 + num_attention_heads: 28 + num_key_value_heads: 4 + hidden_act: silu + max_position_embeddings: 32768 + initializer_range: 0.02 + rms_norm_eps: 1e-6 + use_cache: true + tie_word_embeddings: false + rope_theta: 1000000.0 + use_sliding_window: false + sliding_window: 4096 + max_window_layers: 28 + attention_dropout: 0.0 + mask_token_id: 151666 + pad_token_id: 151643 + bos_token_id: 151643 + eos_token_id: 151643 + torch_dtype: bfloat16 diff --git a/xlm-models/dreamon/configs/model_type/dreamon.yaml b/xlm-models/dreamon/configs/model_type/dreamon.yaml new file mode 100644 index 00000000..b12fe123 --- /dev/null +++ b/xlm-models/dreamon/configs/model_type/dreamon.yaml @@ -0,0 +1,20 @@ +# @package _global_ + +lightning_module: + _target_: xlm.harness.Harness + +predictor: + _target_: dreamon.predictor_dreamon.DreamOnPredictor + +loss: + _target_: dreamon.loss_dreamon.DreamOnLoss + token_reweighting: false + alpha: 0.25 + gamma: 2.0 + time_reweighting: linear + weight_eos: true + max_delete: 64 + +tags: + model_type: dreamon + diff --git a/xlm-models/dreamon/configs/model_type/dreamon_eval.yaml b/xlm-models/dreamon/configs/model_type/dreamon_eval.yaml new file mode 100644 index 00000000..30739476 --- /dev/null +++ b/xlm-models/dreamon/configs/model_type/dreamon_eval.yaml @@ -0,0 +1,10 @@ +# @package _global_ +# Eval / inference: merge Harness + DreamOnPredictor at config root (same pattern as ``dream_eval``). +lightning_module: + _target_: xlm.harness.Harness + +predictor: + _target_: dreamon.predictor_dreamon.DreamOnPredictor + +tags: + model_type: dreamon diff --git a/xlm-models/dreamon/configuration_dreamon.py b/xlm-models/dreamon/configuration_dreamon.py new file mode 100644 index 00000000..ede020f1 --- /dev/null +++ b/xlm-models/dreamon/configuration_dreamon.py @@ -0,0 +1,10 @@ +"""DreamOn variant config (expand token and other DreamOn Hub fields).""" + +from xlm.backbones.dream.configuration_base import DreamConfigBase + + +class DreamOnConfig(DreamConfigBase): + def __init__(self, **kwargs): + expand_token_id = kwargs.pop("expand_token_id", 151667) + super().__init__(**kwargs) + self.expand_token_id = expand_token_id \ No newline at end of file diff --git a/xlm-models/dreamon/datamodule_dreamon.py b/xlm-models/dreamon/datamodule_dreamon.py new file mode 100644 index 00000000..7cfd1748 --- /dev/null +++ b/xlm-models/dreamon/datamodule_dreamon.py @@ -0,0 +1,270 @@ + +from dataclasses import dataclass +import random +from typing import Any, Dict, List, Literal, Mapping, Optional, Sequence, Tuple +import torch +from torch import Tensor + +def compute_position_id_with_mask(attention_mask_1d: Tensor) -> Tensor: + """DreamOn-compatible position_ids from a 1D attention mask.""" + if attention_mask_1d.dim() != 1: + raise ValueError( + f"attention_mask_1d must be 1D, got shape={tuple(attention_mask_1d.shape)}" + ) + pos = torch.cumsum(attention_mask_1d.to(torch.long), dim=0) - 1 + pos = pos.clamp_min(0) + return pos * attention_mask_1d.to(torch.long) + + +def masking_merge_for_response(input_tokens, tokenizer, merge_prob=0.5, merge_schedule="dynamic_inverse", use_uniform_merge_prob=0.5): + """ + The process is: + 1. Independently mask each token with probability sampling_ratio. + If a token is masked it is replaced by "", otherwise it remains unchanged. + 2. Scan the masked sequence for adjacent "" tokens. Whenever found, with probability merge_prob: + - Mark the first token's label as "" (indicating the head of a merged pair). + - Modify the attention_mask so that the second token is not attended to (i.e. set its attention_mask to 0). + Tokens that are not part of a merge or are not masked are labeled as "". + 3. Compute position_ids such that effective tokens (attention_mask==1) receive sequential indices, + while merged-out tokens (attention_mask==0) receive a default position of 0. + + Parameters: + input_tokens (torch.Tensor): The original sequence of tokens as a tensor. + sampling_ratio (float): The independent probability a token is replaced with "". + merge_prob (float): The probability that a pair of adjacent "" tokens are merged. + + Returns: + Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + - final_tokens: The token tensor after independent masking (tokens remain as masked or original). + - labels: A tensor (same length as tokens) with each token labeled as 1 ("") or 0 (""). + - attention_mask: A tensor of binary values (1 for effective tokens, 0 for merged tokens). + """ + sampling_ratio = torch.rand(1) + # Step 1: Sampling masks + mask = torch.rand_like(input_tokens, dtype=torch.float) < sampling_ratio + final_tokens = input_tokens.clone() + final_tokens[mask] = tokenizer.mask_token_id + + eos_mask = input_tokens == tokenizer.eos_token_id + final_tokens[eos_mask] = tokenizer.mask_token_id + + # Initialize labels and attention_mask + labels = input_tokens.clone() + attention_mask = torch.ones_like(input_tokens, dtype=torch.long) + + ## Step 2: Merge + num_masked = mask.sum().item() + + if torch.rand(1).item() < use_uniform_merge_prob: + merge_schedule = "static" + if merge_schedule == "dynamic_inverse": + dynamic_merge_prob = merge_prob * (1 - (num_masked / input_tokens.size(0))) + elif merge_schedule == "dynamic_proportional": + dynamic_merge_prob = merge_prob * (num_masked / input_tokens.size(0)) + elif merge_schedule == "static": + dynamic_merge_prob = merge_prob + elif merge_schedule == "random": + dynamic_merge_prob = torch.rand(1).item() * merge_prob + elif merge_schedule == "full_random": + # So we need to vary merge_prob to [0,1] to make the model more robust + dynamic_merge_prob = torch.rand(1).clamp(0.0, 0.95) + else: + raise ValueError(f"Unknown merge schedule: {merge_schedule}") + + rand_values = torch.rand(len(final_tokens)-1) + + for i in range(len(final_tokens)-1): + if input_tokens[i] == tokenizer.eos_token_id: + break + if (final_tokens[i] == tokenizer.mask_token_id and + final_tokens[i+1] == tokenizer.mask_token_id and ## adjacement MASK + rand_values[i] < dynamic_merge_prob): ## merge + labels[i] = tokenizer.expand_token_id + attention_mask[i+1] = 0 + + + return final_tokens, labels, attention_mask, sampling_ratio + +def get_split_ids(example,tokenizer,middle_strategy,middle_line_num): + if middle_strategy == 'random': + response_ids = example['input_ids'] + # Split response into prefix, middle, suffix + total_length = len(response_ids) + mid_start = random.randint(1, total_length - 2) + mid_end = random.randint(mid_start + 1, total_length - 1) + prefix_ids = response_ids[:mid_start] + middle_ids = response_ids[mid_start:mid_end] + suffix_ids = response_ids[mid_end:] + elif middle_strategy == 'line': + # Split response into prefix, middle, suffix — now using line-based selection + prefix_str, code_block, suffix_str = example['prefix'], example['middle'], example['suffix'] + code_lines = code_block.split('\n') + # Choose start and end indices for middle (consecutive lines) + max_attempts = 5 + for _ in range(max_attempts): + # Ensure valid random indices + if not middle_line_num: + try: + middle_start = random.randint(1, len(code_lines) - 2) + middle_end = random.randint(middle_start + 1, len(code_lines) - 1) + except: + middle_start = 1 + middle_end = len(code_lines) + else: + try: + middle_start = random.randint(1, len(code_lines) - middle_line_num - 1) + middle_end = middle_start + middle_line_num + except: + middle_start = 1 + middle_end = len(code_lines) + + # Extract the slice + middle_lines = code_lines[middle_start:middle_end] + middle_str = "\n".join(middle_lines) + '\n' + + # Check your required conditions + if len(middle_str.split()) >= 3 and 'def' not in middle_str and len(middle_str) > 3: + break # Exit loop early if valid one is found + + + prefix_lines = code_lines[:middle_start] + suffix_lines = code_lines[middle_end:] + + # Join the lines back into strings + prefix_str = prefix_str + "\n".join(prefix_lines) + '\n' + suffix_str = "\n".join(suffix_lines) + suffix_str + + prefix_ids = tokenizer.encode(prefix_str) + middle_ids = tokenizer.encode(middle_str) + suffix_ids = tokenizer.encode(suffix_str) + + return torch.tensor(prefix_ids), torch.tensor(middle_ids), torch.tensor(suffix_ids) + + +class DreamOnInfillTrainCollator: + + def __init__( + self, + tokenizer: Any, + max_length: int = 1024, + truncation: Literal["error", "left", "right"] = "error", + middle_strategy: Literal["line", "random"] = "line", + middle_line_num: Optional[int] = None, + merge_prob: float = 0.5, + max_delete: int = 64, + merge_schedule: str = "dynamic_inverse", + use_uniform_merge_prob: float = 0.5, + expand_token_id: int = 151667, + ): + self.tokenizer = tokenizer + self.max_length = max_length + self.truncation = truncation + self.middle_strategy = middle_strategy + self.middle_line_num = middle_line_num + self.merge_prob = merge_prob + self.max_delete = max_delete + self.merge_schedule = merge_schedule + self.use_uniform_merge_prob = use_uniform_merge_prob + self.expand_token_id = expand_token_id + + def __call__( + self, examples: List[Dict[str, Any]] + ) -> Dict[str, torch.Tensor]: + if self.truncation not in ("error", "left", "right"): + raise ValueError(f"Invalid truncation={self.truncation}") + + input_ids_batch = [] + labels_batch = [] + attention_mask_batch = [] + position_ids_batch = [] + loss_mask_batch = [] + t_batch = [] + for e in examples: + prefix_ids, middle_ids, suffix_ids = get_split_ids(e, self.tokenizer, self.middle_strategy, self.middle_line_num) + prompt_ids = torch.tensor(e['prompt_ids']) + prompt_length = prompt_ids.shape[0] + response_length = prefix_ids.shape[0] + middle_ids.shape[0] + suffix_ids.shape[0] + + # EOS token handling + if self.max_length - prompt_length - response_length > 0 and self.max_delete > 0: + eos_count = torch.randint( + low=0, + high=min(self.max_delete, self.max_length - prompt_length - response_length), + size=(1,), + ).item() + eos_tensor = torch.tensor([self.tokenizer.eos_token_id] * eos_count, dtype=middle_ids.dtype) + else: + eos_count = 0 + eos_tensor = torch.tensor([], dtype=middle_ids.dtype) + + middle_ids = torch.cat([middle_ids, eos_tensor]) + + masked_middle_ids, labels, middle_attention_mask, t = masking_merge_for_response( + middle_ids, + self.tokenizer, + merge_prob=self.merge_prob, + merge_schedule=self.merge_schedule, + use_uniform_merge_prob=self.use_uniform_merge_prob + ) + # Concat all parts + input_ids = torch.cat([ + prompt_ids, + prefix_ids, + masked_middle_ids, + suffix_ids + ], dim=-1) + + attention_mask = torch.cat([ + torch.ones_like(prompt_ids), + torch.ones_like(prefix_ids), + middle_attention_mask, + torch.ones_like(suffix_ids) + ], dim=-1) + + labels = torch.cat([ + prompt_ids, + prefix_ids, + labels, + suffix_ids + ], dim = -1) + + # Padding or Truncation + sequence_length = input_ids.shape[0] + if sequence_length < self.max_length: + pad_len = self.max_length - sequence_length + input_ids = torch.cat([input_ids, torch.full((pad_len,), self.tokenizer.pad_token_id, dtype=input_ids.dtype)]) + labels = torch.cat([labels, torch.full((pad_len,), self.tokenizer.pad_token_id, dtype = labels.dtype)]) + attention_mask = torch.cat([attention_mask, torch.ones(pad_len, dtype=attention_mask.dtype)]) + elif sequence_length > self.max_length: + if self.truncation == "left": + input_ids = input_ids[-self.max_length:] + labels = labels[-self.max_length:] + attention_mask = attention_mask[-self.max_length:] + elif self.truncation == "right": + input_ids = input_ids[:self.max_length] + labels = labels[:self.max_length] + attention_mask = attention_mask[:self.max_length] + else: + raise ValueError(f"Unknown truncation strategy: {self.truncation}") + + #position_ids = compute_position_id_with_mask(input_ids != self.tokenizer.pad_token_id) + position_ids = compute_position_id_with_mask(attention_mask) + + # Loss mask (only for merged part) + loss_mask = (input_ids == self.tokenizer.mask_token_id) & (attention_mask == 1) + + input_ids_batch.append(input_ids) + labels_batch.append(labels) + attention_mask_batch.append(attention_mask) + position_ids_batch.append(position_ids) + loss_mask_batch.append(loss_mask) + t_batch.append(t) + + return { + "input_ids": torch.stack(input_ids_batch, dim=0), + "labels": torch.stack(labels_batch, dim=0), + "attention_mask": torch.stack(attention_mask_batch, dim=0), + "position_ids": torch.stack(position_ids_batch, dim=0), + "loss_mask": torch.stack(loss_mask_batch, dim=0), + "t": torch.stack(t_batch, dim=0) + } diff --git a/xlm-models/dreamon/dreamon_model.py b/xlm-models/dreamon/dreamon_model.py new file mode 100644 index 00000000..9524b81f --- /dev/null +++ b/xlm-models/dreamon/dreamon_model.py @@ -0,0 +1,25 @@ +"""DreamOn model: backbone + variable-canvas diffusion generation.""" + + +from transformers import PretrainedConfig +from xlm.backbones.dream.modeling_dream import DreamModelCore +from .configuration_dreamon import DreamOnConfig + + +class DreamOnModel(DreamModelCore): + """DreamOn decoder with expand/delete `diffusion_generate`.""" + + config_class = DreamOnConfig + + def get_named_params_for_weight_decay(self): + # all parameters except biases and layer-norm parameters + for name, param in self.named_parameters(): + if "bias" in name or "norm" in name: + continue + yield (name, param) + + def get_named_params_for_no_weight_decay(self): + # biases and layer-norm parameters + for name, param in self.named_parameters(): + if "bias" in name or "norm" in name: + yield (name, param) diff --git a/xlm-models/dreamon/loss_dreamon.py b/xlm-models/dreamon/loss_dreamon.py new file mode 100644 index 00000000..77578435 --- /dev/null +++ b/xlm-models/dreamon/loss_dreamon.py @@ -0,0 +1,133 @@ +from dataclasses import dataclass +from typing import Any, Dict, Literal, Optional +import torch +from torch import nn + + +class DreamOnLoss: + + def __init__( + self, + model: Optional[Any] = None, + tokenizer: Optional[Any] = None, + token_reweighting: bool = False, + alpha: float = 0.25, + gamma: float = 2.0, + time_reweighting: Optional[Literal["linear"]] = None, + weight_eos: bool = False, + max_delete: int = 64, + ): + self.model = model + self.tokenizer = tokenizer + self.token_reweighting = token_reweighting + self.alpha = alpha + self.gamma = gamma + self.time_reweighting = time_reweighting + self.weight_eos = weight_eos + self.max_delete = max_delete + + def __call__( + self, + batch, + batch_idx: Optional[int] = None, + dataloader_idx: Optional[int] = None, + dataloader_name: Optional[str] = None, + ): + loss_dict = self.loss_fn( + batch, batch_idx, dataloader_idx, dataloader_name + ) + return loss_dict + + def loss_fn( + self, + batch: Dict[str, Any], + batch_idx: Optional[int] = None, + dataloader_idx: Optional[int] = None, + dataloader_name: Optional[str] = None, + ) -> Dict[str, Any]: + del batch_idx, dataloader_idx, dataloader_name + input_ids = batch["input_ids"].to(self.model.device) + labels = batch["labels"].to(self.model.device) + attention_mask = batch["attention_mask"].to(self.model.device) + position_ids = batch["position_ids"].to(self.model.device) + loss_mask = batch["loss_mask"].to(self.model.device) + t = batch["t"].to(self.model.device) + loss_fct = nn.CrossEntropyLoss(reduction="none") + + if attention_mask.dim() == 2: + # Input is (B, S) -> need to create pairwise mask (B, S, S) + attention_mask = torch.logical_and( + attention_mask.unsqueeze(1).unsqueeze(-2), # (B, 1, S, 1) + attention_mask.unsqueeze(1).unsqueeze(-1) # (B, 1, S, 1) + ) # Result: (B, 1, S, S) + + elif attention_mask.dim() == 3: + # Already (B, S, S), just add head dimension + attention_mask = attention_mask.unsqueeze(1) # (B, 1, S, S) + else: + raise ValueError(f"Unsupported attention_mask shape: {attention_mask.shape}") + + loss_mask = loss_mask.reshape(-1) + + output = self.model( + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + use_cache=False, + ) + logits = output.logits + + shift_logits = torch.cat( + [logits[:, 0:1], logits[:, :-1]], dim=1 + ).contiguous() + shift_labels = labels.contiguous() + # Flatten the tokens + shift_logits = shift_logits.view(-1, self.model.config.vocab_size) + shift_labels = shift_labels.view(-1) + # Enable model parallelism + shift_labels = shift_labels.to(shift_logits.device) + loss = loss_fct(shift_logits, shift_labels) + + # We use weighted loss + loss_mask = loss_mask.to(loss.device) + loss = loss.masked_fill(~loss_mask, 0) + if self.token_reweighting: + loss = ( + self.alpha + * (1 - torch.exp(-loss)) ** self.gamma + * loss + ) + + if self.time_reweighting == "original": + raise NotImplementedError + weight = 1 / t[:, None].float().expand(labels.size()) + elif self.time_reweighting == "linear": + weight = 1 - t.float().expand(labels.size()) + else: + raise NotImplementedError + + loss = loss * weight.reshape(-1) + + if self.weight_eos and self.max_delete > 0: + non_eos_mask = (shift_labels != self.tokenizer.eos_token_id) & loss_mask + non_eos_loss = loss.clone() + non_eos_loss[~non_eos_mask] = 0 + non_eos_count = non_eos_mask.sum().item() + non_eos_loss = non_eos_loss.sum() + + + eos_mask = (shift_labels == self.tokenizer.eos_token_id) & loss_mask + eos_loss = loss.clone() + eos_loss[~eos_mask] = 0 + eos_count = eos_mask.sum().item() + eos_loss = eos_loss.sum() / eos_count + + + loss = (non_eos_loss + eos_loss) / (non_eos_count + 1) + else: + valid_token_this_rank = torch.sum(loss_mask) + + + loss = torch.sum(loss) / valid_token_this_rank + + return {"loss": loss} diff --git a/xlm-models/dreamon/predictor_dreamon.py b/xlm-models/dreamon/predictor_dreamon.py new file mode 100644 index 00000000..77265cdd --- /dev/null +++ b/xlm-models/dreamon/predictor_dreamon.py @@ -0,0 +1,468 @@ +from typing import Any, Dict, List, Optional + +import torch +import torch.nn.functional as F +from xlm.datamodule import Tokenizer +from xlm.harness import Predictor +from xlm.noise import NoiseSchedule +import torch.distributions as dists +from .dreamon_model import DreamOnModel + +def top_p_logits(logits, top_p=None): + sorted_logits, sorted_indices = torch.sort(logits, descending=True) + cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1) + sorted_indices_to_remove = cumulative_probs > top_p + # Shift the indices to the right to keep the first token above the threshold + sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone() + sorted_indices_to_remove[..., 0] = 0 + + mask = torch.zeros_like(logits, dtype=torch.bool, device=logits.device) + mask = mask.scatter_(-1, sorted_indices, sorted_indices_to_remove) + logits = logits.masked_fill(mask, torch.finfo(logits.dtype).min) + return logits + +def top_k_logits(logits, top_k=None): + top_k = min(top_k, logits.size(-1)) # Safety check + # Remove all tokens with a probability less than the last token of the top-k + indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None] + logits = logits.masked_fill(indices_to_remove, torch.finfo(logits.dtype).min) + return logits + +def sample_tokens(logits, temperature=0.0, top_p=None, top_k=None, margin_confidence=False, neg_entropy=False): + + if temperature > 0: + logits = logits / temperature + if top_p is not None and top_p < 1: + logits = top_p_logits(logits, top_p) + if top_k is not None: + logits = top_k_logits(logits, top_k) + probs = torch.softmax(logits, dim=-1) + + if temperature > 0: + try: + x0 = dists.Categorical(probs=probs).sample() + confidence = torch.gather(probs, -1, x0.unsqueeze(-1)).squeeze(-1) + except: + confidence, x0 = probs.max(dim=-1) + else: + confidence, x0 = probs.max(dim=-1) + + if margin_confidence: + sorted_probs, _ = torch.sort(probs, dim=-1, descending=True) + # Extract top1 and top2 probabilities + top1_probs = sorted_probs[:, 0] + top2_probs = sorted_probs[:, 1] + # Calculate confidence as top1 - top2 + confidence = top1_probs - top2_probs + + if neg_entropy: + epsilon = 1e-10 + log_probs = torch.log(probs + epsilon) + confidence = torch.sum(probs * log_probs, dim=-1) + + return confidence, x0 + +def _dreamon_mdm_batch_generate( + model: DreamOnModel, + x: torch.LongTensor, + *, + mask_id: int, + expand_id: int, + steps: int, + eps: float, + temperature: float, + top_p: Optional[float], + top_k: Optional[int], + alg: str, + alg_temp: Optional[float], + show_progress: bool = False, + decode_fn: Optional[Any] = None, +) -> torch.LongTensor: + """Port of ``MDMGenerator.batch_generate`` using a Dream/DreamOn ``model`` forward.""" + device = x.device + x = x.clone() + timesteps = torch.linspace(1, eps, steps + 1, device=device) + for i in range(steps): + mask_index = x == mask_id + logits = model(x).logits + logits = torch.cat([logits[:, :1], logits[:, :-1]], dim=1) + logits = logits[mask_index] + if ( + expand_id is not None + and expand_id >= 0 + and expand_id < logits.shape[-1] + ): + logits[:, expand_id] -= 1e9 + t = timesteps[i] + s = timesteps[i + 1] + if torch.all(~mask_index): + break + + if alg == "origin": + p_transfer = 1 - s / t if i < steps - 1 else 1 + x0 = torch.zeros_like(x[mask_index], device=device, dtype=torch.long) + mask_id + transfer_index_t_s = torch.rand(*x0.shape, device=device) < p_transfer + _, x0[transfer_index_t_s] = sample_tokens( + logits[transfer_index_t_s], + temperature=temperature, + top_p=top_p, + top_k=top_k, + ) + x[mask_index] = x0.clone() + else: + if alg == "maskgit_plus": + confidence, x0 = sample_tokens( + logits, temperature=temperature, top_p=top_p, top_k=top_k + ) + elif alg == "topk_margin": + confidence, x0 = sample_tokens( + logits, + temperature=temperature, + top_p=top_p, + top_k=top_k, + margin_confidence=True, + ) + elif alg == "entropy": + confidence, x0 = sample_tokens( + logits, + temperature=temperature, + top_p=top_p, + top_k=top_k, + neg_entropy=True, + ) + else: + raise RuntimeError(f"Unknown alg: {alg}") + + number_transfer_tokens = 1 + if number_transfer_tokens > 0: + if alg_temp is None or alg_temp == 0: + _, transfer_index = torch.topk(confidence, number_transfer_tokens) + else: + confidence = confidence / alg_temp + confidence = F.softmax(confidence, dim=-1) + transfer_index = torch.multinomial( + confidence, num_samples=number_transfer_tokens + ) + x0_ = torch.zeros_like(x0, device=device, dtype=torch.long) + mask_id + x0_[transfer_index] = x0[transfer_index].clone() + x[mask_index] = x0_ + + if show_progress and decode_fn is not None: + print("==" * 50 + f" step {i} " + "==" * 50) + print(decode_fn(x[0].tolist())) + + return x + + +def _dreamon_mdm_batch_generate_with_expand( + model: DreamOnModel, + input_ids: torch.LongTensor, + *, + mask_id: int, + expand_id: int, + eos_id: int, + max_tokens: int, + max_gen_len: int, + min_gen_len: int, + steps: int, + eps: float, + temperature: float, + top_p: Optional[float], + top_k: Optional[int], + alg: str, + alg_temp: Optional[float], + number_transfer_tokens: int, + expand_budget: int, + pad_eos_to_right: bool, + delete_eos_token: bool, + show_progress: bool = False, + decode_fn: Optional[Any] = None, +) -> torch.LongTensor: + """Port of ``MDMGenerator.batch_generate_with_expand_as_token`` for DreamOn.""" + device = input_ids.device + max_tokens = min( + max_tokens, input_ids.shape[1] + max_gen_len - min_gen_len + ) + x = F.pad( + input_ids, + (0, max_tokens - input_ids.shape[1]), + value=mask_id, + ) + num_generation_tokens = min_gen_len + expand_budget_local = expand_budget + + for i in range(steps): + cur_generation_window_length = ( + input_ids.shape[1] - min_gen_len + num_generation_tokens + ) + attention_mask = torch.ones( + [input_ids.shape[0], cur_generation_window_length], + dtype=torch.int16, + device=device, + ) + attention_mask = F.pad( + attention_mask, + (0, max_tokens - attention_mask.shape[1]), + value=0, + ) + + mask_index = (x == mask_id) & (attention_mask == 1) + if torch.all(~mask_index[:, :cur_generation_window_length]): + break + + tok_idx = attention_mask.long().cumsum(-1) - 1 + tok_idx.masked_fill_(attention_mask == 0, 1) + + attention_mask_2d = torch.logical_and( + attention_mask.unsqueeze(1).unsqueeze(-2), + attention_mask.unsqueeze(1).unsqueeze(-1), + ) + + output = model(x, attention_mask_2d, tok_idx) + logits = output.logits + logits = torch.cat([logits[:, :1], logits[:, :-1]], dim=1) + logits = logits[mask_index] + + if cur_generation_window_length == max_tokens or expand_budget_local == 0: + if 0 <= expand_id < logits.shape[-1]: + logits[:, expand_id] -= 1e9 + + if alg == "origin": + raise NotImplementedError("alg='origin' is not supported for mask expansion") + if alg == "maskgit_plus": + confidence, x0 = sample_tokens( + logits, temperature=temperature, top_p=top_p, top_k=top_k + ) + elif alg == "topk_margin": + confidence, x0 = sample_tokens( + logits, + temperature=temperature, + top_p=top_p, + top_k=top_k, + margin_confidence=True, + ) + elif alg == "entropy": + confidence, x0 = sample_tokens( + logits, + temperature=temperature, + top_p=top_p, + top_k=top_k, + neg_entropy=True, + ) + else: + raise RuntimeError(f"Unknown alg: {alg}") + + if number_transfer_tokens > 0: + if alg_temp is None or alg_temp == 0: + _, transfer_index = torch.topk(confidence, number_transfer_tokens) + else: + confidence = confidence / alg_temp + confidence = F.softmax(confidence, dim=-1) + transfer_index = torch.multinomial( + confidence, num_samples=number_transfer_tokens + ) + x0_ = torch.zeros_like(x0, device=device, dtype=torch.long) + mask_id + x0_[transfer_index] = x0[transfer_index].clone() + x[mask_index] = x0_ + + if pad_eos_to_right: + if x.shape[0] != 1: + raise NotImplementedError( + "pad_eos_to_right=True requires batch size 1 (MDMGenerator parity)" + ) + x_seq = x[0] + eos_indices = (x_seq == eos_id).nonzero(as_tuple=True) + if len(eos_indices[0]) > 0: + first_eos_idx = eos_indices[0][0].item() + position_mask = torch.arange(x_seq.size(0), device=device) >= first_eos_idx + replace_mask = position_mask & mask_index[0] + x_seq.masked_fill_(replace_mask, eos_id) + x = x_seq.unsqueeze(0) + + if show_progress and decode_fn is not None: + print("=" * 10 + f"Step {i}" + "=" * 10) + print(decode_fn(x[0, :cur_generation_window_length].tolist())) + + expand_indices = (x[0] == expand_id).nonzero(as_tuple=False).squeeze(1) + if expand_indices.numel() > 0: + for idx in sorted(expand_indices.tolist(), reverse=True): + x = torch.cat( + ( + x[:, :idx], + torch.tensor([[mask_id, mask_id]], device=device), + x[:, idx + 1 :], + ), + dim=1, + ) + num_generation_tokens += 1 + expand_budget_local -= 1 + if x.shape[1] > max_tokens: + x = x[:, :max_tokens] + + if delete_eos_token: + eos_indices = ((x[0] == eos_id) & (mask_index[0] == 1)).nonzero( + as_tuple=False + ).squeeze(1) + if len(eos_indices) > 0 and show_progress: + print("delete token") + for idx in sorted(eos_indices.tolist(), reverse=True): + x = torch.cat( + ( + x[:, :idx], + x[:, idx + 1 :], + torch.tensor([[mask_id]], device=device), + ), + dim=1, + ) + num_generation_tokens -= 1 + + return x, num_generation_tokens + + +class DreamOnPredictor(torch.nn.Module, Predictor[Any, Dict[str, Any]]): + """Harness wrapper: MDM-style fixed canvas vs expand canvas, aligned with humaneval MDMGenerator.""" + + def __init__( + self, + model: Optional[DreamOnModel] = None, + tokenizer: Optional[Tokenizer] = None, + noise_schedule: Optional[NoiseSchedule] = None, + diffusion_kwargs: Optional[Dict[str, Any]] = None, + mask_expansion: bool = False, + ): + """``model`` / ``tokenizer`` default to None so Hydra can instantiate before Harness assigns them.""" + super().__init__() + self.model = model + self.tokenizer = tokenizer + self.noise_schedule = noise_schedule + self.diffusion_kwargs = diffusion_kwargs or {} + self.mask_expansion = mask_expansion + + def _mdm_runtime_cfg(self) -> Dict[str, Any]: + """Defaults mirror ``MDMGeneratorArgs`` / ``DreamOnGenerationConfig``; override via ``diffusion_kwargs``.""" + g = self.diffusion_kwargs + mc = ( + getattr(self.model, "generation_config", None) + if self.model is not None + else None + ) + return { + "steps": int(g.get("steps", getattr(mc, "steps", 512) if mc else 512)), + "eps": float(g.get("eps", getattr(mc, "eps", 1e-3) if mc else 1e-3)), + "temperature": float(g.get("temperature", 0.0)), + "top_p": g.get("top_p", getattr(mc, "top_p", None) if mc else None), + "top_k": g.get("top_k", getattr(mc, "top_k", None) if mc else None), + "alg": str(g.get("alg", getattr(mc, "alg", "entropy") if mc else "entropy")), + "alg_temp": g.get("alg_temp", getattr(mc, "alg_temp", None) if mc else None), + "show_progress": bool(g.get("show_progress", False)), + "max_tokens": int( + g.get("max_tokens", getattr(mc, "max_length", 2048) if mc else 2048) + ), + "max_gen_len": int( + g.get("max_gen_len", getattr(mc, "max_new_tokens", 512) if mc else 512) + ), + "min_gen_len": int( + g.get("min_gen_len", getattr(mc, "min_gen_len", 16) if mc else 16) + ), + "expand_budget": int( + g.get("max_gen_len", getattr(mc, "max_new_tokens", 512) if mc else 512) + ), + "number_transfer_tokens": int( + g.get( + "number_transfer_tokens", + getattr(mc, "number_transfer_tokens", 1) if mc else 1, + ) + ), + "pad_eos_to_right": bool(g.get("pad_eos_to_right", False)), + "delete_eos_token": bool(g.get("delete_eos_token", False)), + } + + @torch.inference_mode + @torch.no_grad() + @torch._dynamo.disable() + def predict( + self, + batch: Any, + batch_idx: Optional[int] = None, + dataloader_idx: Optional[int] = None, + dataloader_name: Optional[str] = None, + ) -> Dict[str, Any]: + del batch_idx, dataloader_idx, dataloader_name + if self.model is None or self.tokenizer is None: + raise RuntimeError( + "DreamOnPredictor.model and .tokenizer must be set (Harness does this after Hydra instantiate)." + ) + input_ids = batch["input_ids"].to(self.model.device) + cfg = self._mdm_runtime_cfg() + mask_id = int(self.tokenizer.mask_token_id) + eos_id = int(self.tokenizer.eos_token_id) + expand_id = int(self.tokenizer.expand_id) + + decode_fn = ( + (lambda ids: self.tokenizer.decode(ids, skip_special_tokens=True)) + if cfg["show_progress"] + else None + ) + prefix_lens = batch["prefix_lens"] + generations = [] + if not self.mask_expansion: + out = _dreamon_mdm_batch_generate( + self.model, + input_ids, + mask_id=mask_id, + expand_id=expand_id, + steps=cfg["steps"], + eps=cfg["eps"], + temperature=cfg["temperature"], + top_p=cfg["top_p"], + top_k=cfg["top_k"], + alg=cfg["alg"], + alg_temp=cfg["alg_temp"], + show_progress=cfg["show_progress"], + decode_fn=decode_fn, + ) + generations.extend([self.tokenizer.decode(g[pl:pl+ml].tolist(), skip_special_tokens = True) for pl, ml, g in zip(prefix_lens, batch["middle_lens"], out)]) + + else: + out, num_generation_tokens = _dreamon_mdm_batch_generate_with_expand( + self.model, + input_ids, + mask_id=mask_id, + expand_id=expand_id, + eos_id=eos_id, + max_tokens=cfg["max_tokens"], + max_gen_len=cfg["max_gen_len"], + min_gen_len=cfg["min_gen_len"], + steps=cfg["steps"], + eps=cfg["eps"], + temperature=cfg["temperature"], + top_p=cfg["top_p"], + top_k=cfg["top_k"], + alg=cfg["alg"], + alg_temp=cfg["alg_temp"], + number_transfer_tokens=cfg["number_transfer_tokens"], + expand_budget=cfg["expand_budget"], + pad_eos_to_right=cfg["pad_eos_to_right"], + delete_eos_token=cfg["delete_eos_token"], + show_progress=cfg["show_progress"], + decode_fn=decode_fn, + ) + generations.append(self.tokenizer.decode(out[0,prefix_lens[0]:prefix_lens[0] + num_generation_tokens].tolist(), skip_special_tokens = True)) + + return {"text": generations} + + def to_dict( + self, + batch: Any, + preds: Dict[str, Any], + batch_idx: Optional[int] = None, + dataloader_idx: Optional[int] = None, + dataloader_name: Optional[str] = None, + ) -> List[Dict[str, Any]]: + del batch, batch_idx, dataloader_idx, dataloader_name + return [{"text": t} for t in preds["text"]] + + def generate(self, prompts: List[str]) -> List[str]: + raise NotImplementedError( + "DreamOnPredictor.generate is not implemented; use predict()." + ) diff --git a/xlm-models/mlm/datamodule_mlm.py b/xlm-models/mlm/datamodule_mlm.py index 06fce899..b9321a8e 100644 --- a/xlm-models/mlm/datamodule_mlm.py +++ b/xlm-models/mlm/datamodule_mlm.py @@ -162,6 +162,7 @@ def prepare_prefix_ids( pad_token_id: int, max_seq_len: Optional[int] = None, truncate: Literal["max", "block", None] = "block", + pad_left: bool = True ) -> Dict[str, TT]: """ Prepare prefix ids for seq2seq tasks. @@ -200,11 +201,11 @@ def prepare_prefix_ids( temp, max_len, pad_token_id, - pad_left=True, + pad_left=pad_left, ) ) attention_mask.append( - pad_truncate_list([1] * len(temp), max_len, 0, pad_left=True) + pad_truncate_list([1] * len(temp), max_len, 0, pad_left=pad_left) ) return { @@ -659,6 +660,85 @@ def __call__( ) return batch +class MLMDreamInfillPredCollator(Collator): + def __init__( + self, + tokenizer: Tokenizer, + add_bos: bool = True, + add_eos: bool = True, + pass_through_fields: Optional[List[str]] = None, + mask_expansion: bool = False, + fix_middle_length: Optional[int] = None, + max_tokens: Optional[int] = None, + max_prompt_len: Optional[int] = None, + max_gen_len: Optional[int] = None, + min_gen_len: Optional[int] = None, + pad_to_max_len: bool = False, + ): + self.tokenizer = tokenizer + self.add_bos = add_bos + self.add_eos = add_eos + self.pass_through_fields = ( + list(pass_through_fields) + if pass_through_fields is not None + else [] + ) + self.mask_expansion = mask_expansion + self.fix_middle_length = fix_middle_length + self.max_tokens = max_tokens + self.max_prompt_len = max_prompt_len + self.max_gen_len = max_gen_len + self.min_gen_len = min_gen_len + self.pad_to_max_len = pad_to_max_len + + def __call__( + self, + examples: List[BaseCollatorInput], + ) -> MLMBatch: + def _middle_mask_len(e: Mapping[str, Any]) -> int: + if self.fix_middle_length is not None: + return int(self.fix_middle_length) + if not self.mask_expansion: + return len(e["middle_ids"]) + return self.min_gen_len + + input_ids = [ + [self.tokenizer.bos_token_id] * int(self.add_bos) + + list(e["prefix_ids"]) + + [self.tokenizer.mask_token_id] * _middle_mask_len(e) + + list(e['suffix_ids']) + + [self.tokenizer.eos_token_id] * int(self.add_eos) + for e in examples + ] + + input_ids = [p[-self.max_prompt_len:] for p in input_ids] + if not self.mask_expansion: + max_seq_len = self.max_tokens if self.pad_to_max_len else min(self.max_tokens, max([len(p)+self.max_gen_len for p in input_ids])) + batch = prepare_prefix_ids( + input_ids, + self.tokenizer.pad_token_id, + max_seq_len=max_seq_len, + truncate="block", + pad_left=False + ) + else: # Keep batch size as 1 + max_prompt_len_batch = max(len(p) for p in input_ids) + batch = prepare_prefix_ids( + input_ids, + self.tokenizer.pad_token_id, + max_seq_len=max_prompt_len_batch, + truncate="block", + pad_left=False + ) + batch["prefix_lens"] = [len([self.tokenizer.bos_token_id] * int(self.add_bos) + + list(e["prefix_ids"])) for e in examples] + batch["middle_lens"] = [len(e["middle_ids"]) for e in examples] + + for key in self.pass_through_fields: + if key in examples[0]: + batch[key] = [ex[key] for ex in examples] + return batch + def _replace_100_with_pad(ids: torch.Tensor, tokenizer: Tokenizer): _ids = ids.clone() diff --git a/xlm-models/xlm_models.json b/xlm-models/xlm_models.json index 244c133d..16948c3f 100644 --- a/xlm-models/xlm_models.json +++ b/xlm-models/xlm_models.json @@ -4,5 +4,6 @@ "mlm": "mlm", "mdlm": "mdlm", "flexmdm": "flexmdm", - "dream": "dream" + "dream": "dream", + "dreamon": "dreamon" } \ No newline at end of file