CoolFace
Apppublic

chendl/compositional_test

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
consolidate_rag_checkpoint.py102 linesDownload Raw Back to rag
1"""2A script creating a RAG checkpoint from a generator and a question encoder checkpoints.3"""4 5import argparse6from pathlib import Path7 8from transformers import AutoConfig, AutoTokenizer, RagConfig, RagSequenceForGeneration, RagTokenForGeneration9 10 11def consolidate(12    model_type,13    generator_name_or_path: str,14    question_encoder_name_or_path: str,15    dest_dir: Path,16    config_name_or_path: str = None,17    generator_tokenizer_name_or_path: str = None,18    question_encoder_tokenizer_name_or_path: str = None,19):20    if config_name_or_path is None:21        config_name_or_path = "facebook/rag-token-base" if model_type == "rag_token" else "facebook/rag-sequence-base"22 23    if generator_tokenizer_name_or_path is None:24        generator_tokenizer_name_or_path = generator_name_or_path25 26    if question_encoder_tokenizer_name_or_path is None:27        question_encoder_tokenizer_name_or_path = question_encoder_name_or_path28 29    model_class = RagTokenForGeneration if model_type == "rag_token" else RagSequenceForGeneration30 31    # Save model.32    rag_config = RagConfig.from_pretrained(config_name_or_path)33    gen_config = AutoConfig.from_pretrained(generator_name_or_path)34    question_encoder_config = AutoConfig.from_pretrained(question_encoder_name_or_path)35 36    rag_config.generator = gen_config37    rag_config.question_encoder = question_encoder_config38 39    rag_model = model_class.from_pretrained_question_encoder_generator(40        question_encoder_name_or_path, generator_name_or_path, config=rag_config41    )42    rag_model.save_pretrained(dest_dir)43 44    # Sanity check.45    model_class.from_pretrained(dest_dir)46 47    # Save tokenizers.48    gen_tokenizer = AutoTokenizer.from_pretrained(generator_tokenizer_name_or_path)49    gen_tokenizer.save_pretrained(dest_dir / "generator_tokenizer/")50    question_encoder_tokenizer = AutoTokenizer.from_pretrained(question_encoder_tokenizer_name_or_path)51    question_encoder_tokenizer.save_pretrained(dest_dir / "question_encoder_tokenizer/")52 53 54if __name__ == "__main__":55    parser = argparse.ArgumentParser()56    parser.add_argument(57        "--model_type",58        choices=["rag_sequence", "rag_token"],59        required=True,60        type=str,61        help="RAG model type: rag_sequence, rag_token",62    )63    parser.add_argument("--dest", type=str, required=True, help="Path to the output checkpoint directory.")64    parser.add_argument("--generator_name_or_path", type=str, required=True, help="Generator model identifier")65    parser.add_argument(66        "--question_encoder_name_or_path", type=str, required=True, help="Question encoder model identifier"67    )68 69    parser.add_argument(70        "--generator_tokenizer_name_or_path",71        type=str,72        help="Generator tokenizer identifier, if not specified, resolves to ``generator_name_or_path``",73    )74    parser.add_argument(75        "--question_encoder_tokenizer_name_or_path",76        type=str,77        help="Question encoder tokenizer identifier, if not specified, resolves to ``question_encoder_name_or_path``",78    )79    parser.add_argument(80        "--config_name_or_path",81        type=str,82        help=(83            "Identifier of the model config to use, if not provided, resolves to a base config for a given"84            " ``model_type``"85        ),86    )87 88    args = parser.parse_args()89 90    dest_dir = Path(args.dest)91    dest_dir.mkdir(exist_ok=True)92 93    consolidate(94        args.model_type,95        args.generator_name_or_path,96        args.question_encoder_name_or_path,97        dest_dir,98        args.config_name_or_path,99        args.generator_tokenizer_name_or_path,100        args.question_encoder_tokenizer_name_or_path,101    )102