Regresja logistyczna do przewidywania prawdopodobieństw

Nadzorowane uczenie maszynowe w R: regresja

Nina Zumel and John Mount

Win-Vector LLC

Przewidywanie prawdopodobieństw

  • Przewidywanie czy zdarzenie nastąpi (tak/nie): klasyfikacja
  • Przewidywanie prawdopodobieństwa zdarzenia: regresja
  • Regresja liniowa: przewiduje wartości z [$-\infty$, $\infty$]
  • Prawdopodobieństwa: ograniczone do przedziału [0,1]
    • Dlatego stosuje się podejście nieliniowe
Nadzorowane uczenie maszynowe w R: regresja

Przykład: przewidywanie dystrofii mięśniowej Duchenne'a (DMD)

  • wynik: has_dmd    zmienne wejściowe: CK, H
Nadzorowane uczenie maszynowe w R: regresja

Model regresji liniowej

model <- lm(has_dmd ~ CK + H, 
            data = train)

test$pred <- predict(
    model, 
    newdata = test
)

wynik: has_dmd $\in$ {0,1}

  • 0: FALSE
  • 1: TRUE

Model przewiduje wartości spoza zakresu [0:1]

Nadzorowane uczenie maszynowe w R: regresja

Regresja logistyczna

$$ log(\frac{p}{1-p}) = \beta_0 + \beta_1 x_1 + \beta_2 x_2 + ... $$

glm(formula, data, family = binomial)
  • Uogólniony model liniowy
  • Zakłada addytywność i liniowość zmiennych w log-ilorazie szans: $log( p/(1-p) )$
  • family: opisuje rozkład błędów modelu
    • regresja logistyczna: family = binomial
Nadzorowane uczenie maszynowe w R: regresja

Model DMD

model <- glm(has_dmd ~ CK + H, data = train, family = binomial)
  • wynik: dwie klasy, np. $a$ i $b$
  • model zwraca $Prob(b)$
    • Zalecane: 0/1 lub FALSE/TRUE
Nadzorowane uczenie maszynowe w R: regresja

Interpretacja modeli regresji logistycznej

model
Call:  glm(formula = has_dmd ~ CK + H, family = binomial, data = train)

Coefficients:
(Intercept)           CK            H  
  -16.22046      0.07128      0.12552  

Degrees of Freedom: 86 Total (i.e. Null);  84 Residual
Null Deviance:       110.8 
Residual Deviance: 45.16     AIC: 51.16
Nadzorowane uczenie maszynowe w R: regresja

Przewidywanie za pomocą modelu glm()

predict(model, newdata, type = "response")
  • newdata: domyślnie dane treningowe
  • Aby uzyskać prawdopodobieństwa: użyj type = "response"
    • Domyślnie: zwraca log-ilorazy szans
Nadzorowane uczenie maszynowe w R: regresja

Model DMD

model <- glm(has_dmd ~ CK + H, data = train, family = binomial)
test$pred <- predict(model, newdata = test, type = "response")

Nadzorowane uczenie maszynowe w R: regresja

Ocena modelu regresji logistycznej: pseudo-$R^2$

$$ R^2 = 1 - \frac{RSS}{SS_{Tot}} $$

$$ pseudo R^2 = 1 - \frac{deviance}{null.deviance} $$

  • Dewiancja: analogiczna do wariancji (RSS)
  • Dewiancja zerowa: odpowiada $SS_{Tot}$
  • pseudo R^2: wyjaśniona dewiancja
Nadzorowane uczenie maszynowe w R: regresja

Pseudo-$R^2$ na danych treningowych

Użycie broom::glance()

glance(model) %>% 
  summarize(pR2 = 1 - deviance/null.deviance)
   pseudoR2
1 0.5922402

Użycie sigr::wrapChiSqTest()

wrapChiSqTest(model)
"... pseudo-R2=0.59 ..."
Nadzorowane uczenie maszynowe w R: regresja

Pseudo-$R^2$ na danych testowych

# Test data
test %>% 
  mutate(pred = predict(model, newdata = test, type = "response")) %>%
  wrapChiSqTest("pred", "has_dmd", TRUE)

Argumenty:

  • ramka danych
  • nazwa kolumny z prognozami
  • nazwa kolumny z wynikami
  • wartość docelowa (zdarzenie docelowe)
Nadzorowane uczenie maszynowe w R: regresja

Wykres krzywej zysku

GainCurvePlot(test, "pred","has_dmd", "DMD model on test")

Nadzorowane uczenie maszynowe w R: regresja

Czas na ćwiczenia!

Nadzorowane uczenie maszynowe w R: regresja

Preparing Video For Download...