Шаг градиентного спуска для времени доставки

тема: Градиентный шаг · уровень: средний

Условие

Служба доставки еды обучает линейную модель для оценки времени доставки заказа. Для заказа с расстоянием x километров модель предсказывает время p = w*x + b минут, где w и b — текущие параметры модели.

В журнале некоторые фактические времена доставки отсутствуют и обозначены словом NA. Такие строки не участвуют в вычислении ошибки и обновлении параметров. Пусть N — число строк, в которых фактическое время известно.

Используется функция потерь MSE:

L(w, b) = (1 / N) * sum((w*x_i + b - y_i)^2).

Для одного шага градиентного спуска с коэффициентом обучения eta параметры обновляются одновременно по формулам:

dw = (2 / N) * sum(x_i * (w*x_i + b - y_i))

db = (2 / N) * sum(w*x_i + b - y_i)

w_new = w - eta * dw

b_new = b - eta * db

Требуется вывести новые значения w_new и b_new. Равенства чисел в вычислениях не требуют специальной обработки: все строки с известным фактическим временем учитываются по приведённым формулам.

Формат ввода

В первой строке записаны четыре значения: целое число n, вещественные числа eta, w, b.

В следующих n строках записаны расстояние x и фактическое время доставки y. Значение x — целое число. Вместо значения y может быть записано NA.

Формат вывода

Выведите два числа: w_new и b_new, разделённые пробелом.

Каждое число необходимо вывести ровно с четырьмя знаками после десятичной точки. Округление выполняется до ближайшего значения, а при ровно равных расстояниях до двух соседних значений выбирается значение, большее по модулю.

Ограничения

1 <= n <= 2000.

0.001 <= eta <= 0.1.

-10.0 <= w, b <= 10.0.

0 <= x <= 50.

Если значение времени известно, то 1 <= y <= 240.

В записи вещественных чисел eta, w и b не более трёх знаков после десятичной точки.

Гарантируется, что хотя бы в одной строке значение y не равно NA.

Решить задачу с автопроверкой на Python →

Куда дальше