File size: 7,964 Bytes
1f7e277
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
"""Launch the FelaTab x TabArena benchmark on EC2 (replaces the unreliable HF Jobs path).

Pushes the current benchmark scripts to the HF model repo, then boots an on-demand
c7i.8xlarge (32 vCPU / 64 GiB, ~$1.43/hr in us-west-2) whose user-data runs
ec2_run.sh: lite shakedown -> (guard) -> full suite -> upload artifacts to
lowdown-labs/fela-tab-tabarena-results -> poweroff (shutdown behavior = terminate).

Prereqs: `aws login` (SSO session), HF_TOKEN in env (write token for lowdown-labs).

Usage:
  HF_TOKEN=hf_... python launch_ec2.py                      # lite shakedown then full
  HF_TOKEN=hf_... python launch_ec2.py --phases lite        # shakedown only
  HF_TOKEN=hf_... python launch_ec2.py --phases full --datasets "houses,diamonds,..."  # shard

Note: HF_TOKEN is embedded in EC2 user-data (visible to anyone with
DescribeInstanceAttribute permission on this account). Rotate the token after the run
if that matters to you.
"""

import argparse
import os
import time
import urllib.request

import boto3
from huggingface_hub import HfApi, get_token

MODEL_REPO = "lowdown-labs/fela-tab"
INSTANCE_TYPE = "c7i.8xlarge"
KEY_NAME = "felatab-tabarena"
SG_NAME = "felatab-tabarena-ssh"
AMI_PARAM = "/aws/service/canonical/ubuntu/server/24.04/stable/current/amd64/hvm/ebs-gp3/ami-id"

USER_DATA_TMPL = """#!/bin/bash
set -euxo pipefail
export HOME=/root
export DEBIAN_FRONTEND=noninteractive
apt-get update -qq
apt-get install -y -qq git curl python3 python3-venv
curl -LsSf https://astral.sh/uv/install.sh | sh
export PATH="$HOME/.local/bin:$PATH"
uv venv --seed /work-bootvenv
uv pip install --python /work-bootvenv/bin/python -q huggingface_hub
mkdir -p /work
export HF_TOKEN="{hf_token}"
/work-bootvenv/bin/python -c "from huggingface_hub import snapshot_download; snapshot_download('{model_repo}', allow_patterns=['benchmark/tabarena/ec2_run.sh'], local_dir='/work/fela-tab')"
cat > /etc/felatab.env <<EOF
HF_TOKEN={hf_token}
PHASES={phases}
RUN_NAME={run_name}
DATASETS={datasets}
WITH_LGBM={with_lgbm}
SELF_TERMINATE=1
CHECKPOINT={checkpoint}
RESTORE={restore}
EOF
chmod 600 /etc/felatab.env
cat > /etc/systemd/system/felatab-benchmark.service <<'EOF'
[Unit]
Description=FelaTab TabArena benchmark
After=network-online.target
Wants=network-online.target

[Service]
Type=simple
EnvironmentFile=/etc/felatab.env
ExecStart=/bin/bash /work/fela-tab/benchmark/tabarena/ec2_run.sh
User=root
StandardOutput=append:/work/logs_boot.log
StandardError=append:/work/logs_boot.log

[Install]
WantedBy=multi-user.target
EOF
systemctl daemon-reload
systemctl enable --now felatab-benchmark.service
"""


def ensure_key_pair(ec2, key_path):
    names = {k["KeyName"] for k in ec2.describe_key_pairs()["KeyPairs"]}
    if KEY_NAME in names:
        print(f"key pair {KEY_NAME} exists")
        return
    kp = ec2.create_key_pair(KeyName=KEY_NAME)
    with open(key_path, "w") as f:
        f.write(kp["KeyMaterial"])
    os.chmod(key_path, 0o400)
    print(f"created key pair {KEY_NAME} -> {key_path}")


def ensure_security_group(ec2):
    vpcs = ec2.describe_vpcs(Filters=[{"Name": "isDefault", "Values": ["true"]}])["Vpcs"]
    vpc_id = vpcs[0]["VpcId"]
    for sg in ec2.describe_security_groups(
        Filters=[{"Name": "group-name", "Values": [SG_NAME]}, {"Name": "vpc-id", "Values": [vpc_id]}]
    )["SecurityGroups"]:
        print(f"security group {SG_NAME} exists ({sg['GroupId']})")
        return sg["GroupId"]
    my_ip = urllib.request.urlopen("https://checkip.amazonaws.com").read().decode().strip()
    sg = ec2.create_security_group(GroupName=SG_NAME, Description="FelaTab TabArena runner (ssh)", VpcId=vpc_id)
    ec2.authorize_security_group_ingress(
        GroupId=sg["GroupId"],
        IpPermissions=[{"IpProtocol": "tcp", "FromPort": 22, "ToPort": 22,
                        "IpRanges": [{"CidrIp": f"{my_ip}/32", "Description": "launcher ip"}]}],
    )
    print(f"created security group {SG_NAME} ({sg['GroupId']}), ssh from {my_ip}/32")
    return sg["GroupId"]


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--phases", default="lite,full", help="comma list: lite,full or just lite / full")
    ap.add_argument("--datasets", default="", help="optional csv shard (full mode only)")
    ap.add_argument("--run-name", default="felatab_ec2")
    ap.add_argument("--instance-type", default=INSTANCE_TYPE)
    ap.add_argument("--with-lgbm", default="0")
    ap.add_argument("--spot", action="store_true",
                    help="persistent spot with stop-on-interruption: EBS cache survives, "
                         "systemd service resumes the run on restart")
    ap.add_argument("--checkpoint", default="0", choices=["0", "1"],
                    help="upload results cache to the HF repo every 30 min")
    ap.add_argument("--restore", default="0", choices=["0", "1"],
                    help="download existing raw/<run-name> cache before running")
    args = ap.parse_args()

    token = os.environ.get("HF_TOKEN") or get_token()
    assert token, "need HF_TOKEN in env or a logged-in `hf auth login` session"
    api = HfApi(token=token)

    here = os.path.dirname(os.path.abspath(__file__))
    for fname in ("ec2_run.sh", "run_tabarena.py", "fela_ag_model.py"):
        api.upload_file(
            path_or_fileobj=os.path.join(here, fname),
            path_in_repo=f"benchmark/tabarena/{fname}",
            repo_id=MODEL_REPO, repo_type="model",
            commit_message=f"Update TabArena EC2 runner scripts ({fname})",
        )
        print(f"uploaded {fname} -> {MODEL_REPO}")

    ec2 = boto3.client("ec2")
    ssm = boto3.client("ssm")
    ami = ssm.get_parameter(Name=AMI_PARAM)["Parameter"]["Value"]
    print(f"AMI: {ami}")

    key_path = os.path.expanduser(f"~/.aws/{KEY_NAME}.pem")
    ensure_key_pair(ec2, key_path)
    sg_id = ensure_security_group(ec2)

    user_data = USER_DATA_TMPL.format(
        hf_token=token, model_repo=MODEL_REPO, phases=args.phases,
        run_name=args.run_name, datasets=args.datasets, with_lgbm=args.with_lgbm,
        checkpoint=args.checkpoint, restore=args.restore,
    )
    kwargs = dict(
        ImageId=ami,
        InstanceType=args.instance_type,
        KeyName=KEY_NAME,
        MinCount=1, MaxCount=1,
        NetworkInterfaces=[{"DeviceIndex": 0, "Groups": [sg_id], "AssociatePublicIpAddress": True}],
        BlockDeviceMappings=[{"DeviceName": "/dev/sda1",
                              "Ebs": {"VolumeSize": 200, "VolumeType": "gp3", "DeleteOnTermination": True}}],
        TagSpecifications=[{"ResourceType": "instance",
                            "Tags": [{"Key": "Name", "Value": f"felatab-tabarena-{args.run_name}"}]}],
        UserData=user_data,
    )
    if args.spot:
        # persistent spot, stop (not terminate) on interruption: the EBS results cache
        # survives and the systemd service resumes the run when capacity returns
        kwargs["InstanceMarketOptions"] = {
            "MarketType": "spot",
            "SpotOptions": {"SpotInstanceType": "persistent", "InstanceInterruptionBehavior": "stop"},
        }
        kwargs["InstanceInitiatedShutdownBehavior"] = "stop"
    else:
        kwargs["InstanceInitiatedShutdownBehavior"] = "terminate"
    resp = ec2.run_instances(**kwargs)
    iid = resp["Instances"][0]["InstanceId"]
    print(f"instance: {iid} ({args.instance_type}, phases={args.phases}, run={args.run_name})")

    ec2.get_waiter("instance_running").wait(InstanceIds=[iid])
    desc = ec2.describe_instances(InstanceIds=[iid])["Reservations"][0]["Instances"][0]
    ip = desc.get("PublicIpAddress", "<no public ip>")
    print(f"running at {ip}")
    print(f"watch:  ssh -i {key_path} ubuntu@{ip} 'tail -f /work/logs_boot.log'")
    print(f"status: aws ec2 describe-instances --instance-ids {iid} --query 'Reservations[0].Instances[0].State.Name'")
    print("artifacts land in https://huggingface.co/datasets/lowdown-labs/fela-tab-tabarena-results")


if __name__ == "__main__":
    main()