Пока заказов восемь, градиент можно считать по всей базе. При восьми миллионах полный проход ради одного обновления становится непозволительным, и градиент считают по небольшой группе — по батчу.
Так обучают все большие модели без исключения. Но важно понимать вот что: батчевый градиент — не приближение полного, а его оценка. Возьмёшь другую группу — получишь немного другое направление. Сейчас вы увидите, насколько сильно батчи расходятся между собой и что при этом оказывается верно про их среднее.
Подробно: Если данных очень много.
Задание
База выполненных заказов. Три первых числа в строке — признаки заказа: километры, светофоры, свободные курьеры рядом. Четвёртое — сколько заказ ехал на самом деле, в минутах. Это и есть то, что модель учится предсказывать.
База вписана в тест, объявлять её в решении не нужно.
baza = [(4.0, 2, 6, 20.0), (2.0, 5, 4, 19.5), (6.0, 1, 2, 26.5), (1.5, 3, 8, 13.0),
(5.0, 0, 5, 20.5), (3.0, 4, 1, 22.5), (7.5, 2, 3, 32.0), (2.5, 6, 7, 21.0)]
Напишите функцию
def grad_batch(baza, p, ot, do):
# градиент по срезу baza[ot:do]
...
Кода почти не надо: батч — это обычный срез списка, и та же формула производной берёт его без изменений. Делить надо на длину батча, а не всей базы.
Как проверить себя
Посчитайте в точке, где все четыре параметра нулевые: полный градиент по всем восьми заказам, потом градиент по каждому из четырёх батчей — заказы 1–2, 3–4, 5–6, 7–8, — и среднее этих четырёх.
Сравните батчи между собой: по километрам они разойдутся очень сильно. А потом сравните среднее четырёх с полным градиентом. Заметите кое-что: подумайте, всегда ли так будет и что для этого должно быть верно про размеры батчей.
Числа сравниваются с допуском, а не как строки: достаточно совпадения в первых четырёх знаках. Печатать ничего не надо — функция возвращает значение.