Будет ли настоящий Санта использовать нашу модель для выбора подарков?

В 2022 году мы в iterative.ai сделали много замечательных вещей. Одним из них является выпуск MLEM — инструмента с открытым исходным кодом для развертывания модели, если вы его пропустили. И вы, скорее всего, пропустили это, сообщают наши счетчики загрузок.
Итак, этот год MLEMming завершится установкой рождественской елки. И под «установкой рождественской елки» я, конечно же, подразумеваю развертывание модели машинного обучения в облаке. Кажется разумным, чтобы эта модель была DecisionTreeClassifier, которая будет классифицировать имена на «хорошие» и «непослушные» классы.
Получение данных
Чтобы обучить наше дерево, нам нужны обучающие данные, желательно пример реального списка Деда Мороза.
Быстрый поиск в Google показал, что общедоступных образцов нет, однако на официальном веб-сайте Санты есть страница, на которой можно сверить свое имя с этим списком.
Итак, давайте соскрести черта с этого:

Небольшая проверка сети и небольшой анализ DOM позже у нас есть эта функция:
import bs4 as bs4
import requests
def is_nice(name: str):
resp = requests.post("https://www.northpoletimes.com/NaughtyOrNice/", data={"KidName": name, "KidName_submit": "yes"})
resp.raise_for_status()
soup = bs4.BeautifulSoup(resp.text, "html.parser")
div = soup.find("div", attrs={"class": "SantaListKidName"})
if div is None:
return None
img = div.find("img", attrs={"alt": "[Naughty or Nice List]"})
if img is None:
return None
img_src = img.attrs["src"]
return img_src.endswith("NiceCheck.png")
Теперь нам нужен только список имен для проверки. Не будем далеко ходить и воспользуемся первой ссылкой из поиска гугла «самые популярные имена CSV»
Давайте возьмем 1000 лучших имен и проверим, кто из них хороший.
import json
import os.path
from collections import Counter
import pandas as pd
from tqdm import tqdm
def scrape_niceness(names, path):
if os.path.exists(path):
with open(path, "r") as f:
data = json.load(f)
else:
data = {}
for name in tqdm(names):
if name in data:
continue
nice = is_nice(name)
data[name] = nice
with open(path, "w") as f:
json.dump(data, f)
def get_names(count=1000):
data = pd.read_csv("baby-names.csv")
counter = Counter(data.name)
return list(dict(counter.most_common(count)).keys())
names = get_names()
scrape_niceness(names, "nice.json")
Обучение
Пока мы ждем данных, давайте подготовим наш обучающий код. Должно быть довольно просто:
from sklearn.tree import DecisionTreeClassifier
with open("nice.json", "r") as f:
data = json.load(f)
christmas_tree = DecisionTreeClassifier()
christmas_tree.fit(list(data.keys()), list(data.values()))
Сразу же мы столкнулись с проблемой: видимо, имена не являются допустимыми поплавками!
ValueError: could not convert string to float: 'Jesse'
Что ж, давайте объясним нашей машине, что они на самом деле являются числами с плавающей запятой с возможностью встраивания.
Давайте добавим простой этап предварительной обработки в наш рождественский пайплайн, который превратит имена в скрытое состояние модели ALBERT, из которой загружается модель.
from transformers import AlbertTokenizer, AlbertModel
tokenizer = AlbertTokenizer.from_pretrained('albert-base-v2')
model = AlbertModel.from_pretrained("albert-base-v2")
def pre_process(value: str):
encoded_input = tokenizer(value, return_tensors='pt')
output = model(**encoded_input)
return output.last_hidden_state.squeeze(0)[-1].detach().numpy().reshape(1, 768)
Пока мы это делали, нас заблокировал сайт Санты. Кажется, Санта знает, что такое DoS-атака. Но на самом деле это не имеет значения, так как все имена, которые мы проверили, были хорошими (хотя изображение для непослушных имен действительно существует на сервере).
Хорошо, все милы, поэтому мы добавим несколько случайных строк в качестве отрицательных примеров для нашей модели.
Теперь давайте обучим наше дерево и сохраним его с помощью MLEM.
import string
import random
import numpy as np
import mlem
christmas_tree = DecisionTreeClassifier()
with open("nice.json", "r") as f:
data = json.load(f)
data.update(
{
"".join(
random.choice(string.ascii_lowercase)
for _ in range(random.randint(4, 7))
).capitalize(): False
for _ in range(len(data))
}
)
preprocessed = [pre_process(name) for name in tqdm(data)]
christmas_tree.fit(np.stack(preprocessed, axis=1)[0], list(data.values()))
print(christmas_tree.predict(pre_process("Mike"))) # True of course
mdl = mlem.api.save(
christmas_tree,
"christmas_tree",
preprocess=pre_process,
postprocess=post_process,
sample_data="Mike",
)
Мы также добавили небольшой этап постобработки, чтобы получить ответы из бездушных вероятностных чисел.
def post_process(prediction):
if len(prediction.shape) > 1:
return "Nice" if prediction[0][0] < prediction[0][1] else "Naughty"
return "Nice" if prediction[0] else "Naughty"
Бег
Наконец, мы можем опробовать модель локально. MLEM автоматически создаст красивый интерфейс Streamlit, используя sample_data, который мы передали при сохранении модели:
$ mlem serve streamlit -m christmas_tree Starting streamlit server... You can now view your Streamlit app in your browser. URL: <http://0.0.0.0:80>

Хороший!
Чтобы отправить это обратно Деду Морозу, нам нужно где-то его развернуть. MLEM может развертывать модели на ряде платформ, таких как Heroku, Sagemaker и Kubernetes, буквально одной командой.
Так как мы хотели что-то особенное на Рождество, мы экспортировали модель в папку docker build-ready с помощью MLEM и развернули модель на fly.io с помощью flyctl launch — проверьте развертывание!
Это оно! Елку поставили, теперь можно все праздновать!
Счастливого Рождества и счастливых праздников!