CoolFace
Apppublic

philschmid/sagemaker-launcher

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
app.py98 linesDownload Raw Back to root
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