CoolFace
Apppublic

philschmid/sagemaker-launcher

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
trainer.py140 linesDownload Raw Back to root
1from sagemaker.huggingface import HuggingFace2import logging3import sys4from contextlib import contextmanager5from io import StringIO6from streamlit.report_thread import REPORT_CONTEXT_ATTR_NAME7from threading import current_thread8import streamlit as st9import sys10import sagemaker11import boto312 13 14@contextmanager15def st_redirect(src, dst):16    placeholder = st.empty()17    output_func = getattr(placeholder, dst)18 19    with StringIO() as buffer:20        old_write = src.write21 22        def new_write(b):23            if getattr(current_thread(), REPORT_CONTEXT_ATTR_NAME, None):24                buffer.write(b)25                output_func(buffer.getvalue())26            else:27                old_write(b)28 29        try:30            src.write = new_write31            yield32        finally:33            src.write = old_write34 35 36@contextmanager37def st_stdout(dst):38    with st_redirect(sys.stdout, dst):39        yield40 41 42@contextmanager43def st_stderr(dst):44    with st_redirect(sys.stderr, dst):45        yield46 47 48task2script = {49    "text-classification": {50        "entry_point": "run_glue.py",51        "source_dir": "examples/text-classification",52    },53    "token-classification": {54        "entry_point": "run_ner.py",55        "source_dir": "examples/token-classification",56    },57    "question-answering": {58        "entry_point": "run_qa.py",59        "source_dir": "examples/question-answering",60    },61    "summarization": {62        "entry_point": "run_summarization.py",63        "source_dir": "examples/seq2seq",64    },65    "translation": {66        "entry_point": "run_translation.py",67        "source_dir": "examples/seq2seq",68    },69    "causal-language-modeling": {70        "entry_point": "run_clm.py",71        "source_dir": "examples/language-modeling",72    },73    "masked-language-modeling": {74        "entry_point": "run_mlm.py",75        "source_dir": "examples/language-modeling",76    },77}78 79 80def train_estimtator(parameter, config):81    with st_stdout("code"):82        logger = logging.getLogger(__name__)83 84        logging.basicConfig(85            level=logging.getLevelName("INFO"),86            handlers=[logging.StreamHandler(sys.stdout)],87            format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",88        )89        logger.info = print90 91        # git configuration to download our fine-tuning script92        git_config = {"repo": "https://github.com/huggingface/transformers.git", "branch": "v4.4.2"}93 94        # creating fine-tuning script95        entry_point = task2script[parameter["task"]]["entry_point"]96        source_dir = task2script[parameter["task"]]["source_dir"]97        # create train file98        # iam configuration99        session = boto3.session.Session(100            aws_access_key_id=config["aws_access_key_id"],101            aws_secret_access_key=config["aws_secret_accesskey"],102            region_name=config["region"],103        )104        sess = sagemaker.Session(boto_session=session)105 106        iam = session.client(107            "iam", aws_access_key_id=config["aws_access_key_id"], aws_secret_access_key=config["aws_secret_accesskey"]108        )109        role = iam.get_role(RoleName=config["aws_sagemaker_role"])["Role"]["Arn"]110 111        logger.info(f"role: {role}")112        instance_type = config["instance_type"].split("|")[1].split("|")[0].strip()113        logger.info(f"instance_type: {instance_type}")114 115        hyperparameters = {116            "output_dir": "/opt/ml/model",117            "do_train": True,118            "do_eval": True,119            "do_predict": True,120            **parameter,121        }122        del hyperparameters["task"]123        # create estimator124        huggingface_estimator = HuggingFace(125            entry_point=entry_point,126            source_dir=source_dir,127            git_config=git_config,128            base_job_name=config["job_name"],129            instance_type=instance_type,130            sagemaker_session=sess,131            instance_count=config["instance_count"],132            role=role,133            transformers_version="4.4",134            pytorch_version="1.6",135            py_version="py36",136            hyperparameters=hyperparameters,137        )138        # train139        huggingface_estimator.fit()140