philschmid/sagemaker-launcher
0
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 