philschmid/sagemaker-launcher
0
1import streamlit as st2from utils.load_dataset import load_datasets3from utils.load_tasks import load_tasks4from utils.load_models import load_models5from trainer import train_estimtator6from datetime import datetime7import logging8 9logger = logging.getLogger(__name__)10 11 12def main():13 parameter = st.experimental_get_query_params()14 parameter["model_name_or_path"] = parameter.get("model_name_or_path", ["none"])15 parameter["dataset"] = parameter.get("dataset", ["none"])16 parameter["task"] = parameter.get("task", ["none"])17 ### hyperparameter18 parameter["epochs"] = parameter.get("epochs", [3])19 parameter["learning_rate"] = parameter.get("learning_rate", [5e-5])20 parameter["per_device_train_batch_size"] = parameter.get("per_device_train_batch_size", [8])21 parameter["per_device_eval_batch_size"] = parameter.get("per_device_eval_batch_size", [8])22 st.experimental_set_query_params(**parameter)23 24 dataset_list = load_datasets()25 task_list = load_tasks()26 model_list = load_models()27 28 st.header("Hugging Face model & dataset")29 col1, col2 = st.beta_columns(2)30 parameter["model_name_or_path"] = col1.selectbox("Model ID:", parameter["model_name_or_path"] + model_list)31 st.experimental_set_query_params(**parameter)32 33 parameter["dataset"] = col2.selectbox("Dataset:", parameter["dataset"] + dataset_list)34 st.experimental_set_query_params(**parameter)35 36 parameter["task"] = col1.selectbox("Task:", parameter["task"] + task_list)37 st.experimental_set_query_params(**parameter)38 39 use_auth_token = col2.text_input("HF auth token to upload your model:", help="api_xxxxx")40 41 my_expander = st.beta_expander("Hyperparameters")42 col1, col2 = my_expander.beta_columns(2)43 parameter["epochs"] = col1.number_input("Epoch", 3)44 st.experimental_set_query_params(**parameter)45 46 parameter["learning_rate"] = col2.text_input("Learning Rate", 5e-5)47 st.experimental_set_query_params(**parameter)48 49 parameter["per_device_train_batch_size"] = col1.number_input("Training Batch Size", 8)50 st.experimental_set_query_params(**parameter)51 52 parameter["per_device_eval_batch_size"] = col2.number_input("Eval Batch Size", 8)53 st.experimental_set_query_params(**parameter)54 st.markdown("---")55 56 st.header("Amazon Sagemaker configuration")57 58 config = {}59 60 config["job_name"] = st.text_input(61 "model name",62 f"{parameter['model_name_or_path'][0] if isinstance(parameter['model_name_or_path'],list)else parameter['model_name_or_path']}-job-{str(datetime.today()).split()[0]}",63 )64 col1, col2 = st.beta_columns(2)65 66 config["aws_sagemaker_role"] = col1.text_input("AWS IAM role for sagemaker job")67 config["instance_type"] = col2.selectbox(68 "Instance type",69 [70 "single-gpu | ml.p3.2xlarge",71 "multi-gpu | ml.p3.16xlarge",72 ],73 )74 config["region"] = col1.selectbox(75 "AWS Region",76 ["eu-central-1", "eu-west-1", "us-east-1", "us-east-1", "us-west-1", "us-west-2"],77 )78 config["instance_count"] = col2.number_input("Instance count", 1)79 config["use_spot"] = col1.selectbox("use spot instances", [False, True])80 config["distributed"] = col2.selectbox("distributed training", [False, True])81 st.markdown("---")82 83 st.header("Credentials")84 # sagemaker config85 col1, col2 = st.beta_columns(2)86 config["aws_access_key_id"] = col1.text_input("Aws Secret Key ID")87 config["aws_secret_accesskey"] = col2.text_input("Aws Secret Access Key")88 89 if use_auth_token:90 parameter["use_auth_token"] = use_auth_token91 92 if st.button("Start training on SageMaker"):93 train_estimtator(parameter, config)94 95 96if __name__ == "__main__":97 main()98 