Computer >> 컴퓨터 >  >> 프로그래밍 >> Python

Matplotlib에서 statsmodels 선형 회귀(OLS) 결과를 깔끔하게 시각화하는 방법

statsmodels의 선형 회귀(OLS) 분석은 비선형 곡선 형태의 데이터에도 적용할 수 있으며, Matplotlib을 활용하면 회귀 곡선과 신뢰 구간까지 한눈에 들어오는 그래프로 깔끔하게 시각화할 수 있습니다. 이 글에서는 실제 데이터 포인트, 참값(True), OLS 피팅값, 그리고 예측 표준편차 기반의 신뢰 구간을 하나의 그래프에 함께 그리는 방법을 단계별로 살펴보겠습니다.

구현 단계

  1. figure 크기를 설정하고 서브플롯 주변 및 사이의 여백(padding)을 조정합니다.
  2. 결과를 재현할 수 있도록 seed() 메서드로 난수 시드를 고정합니다.
  3. 샘플 수(nsample)와 표준편차(sig) 변수를 초기화합니다.
  4. numpy를 사용하여 x, X, beta, y_true, y 등의 선형 데이터 포인트를 생성합니다.
  5. res는 최소제곱법(Ordinary Least Square, OLS) 클래스의 인스턴스로, sm.OLS(y, X).fit()을 통해 모델을 학습합니다.
  6. wls_prediction_std() 함수로 예측 표준편차와 신뢰 구간의 하한(iv_l), 상한(iv_u)을 계산합니다. 참고로 이 예측 신뢰 구간은 WLS와 OLS에 적용되며, 일반적인 GLS(관측치가 독립적이지만 동일하게 분포하지 않는 경우)에는 적용되지 않습니다.
  7. subplots() 메서드로 figure와 axes 객체를 생성합니다.
  8. plot() 메서드를 사용해 (x, y), (x, y_true), (x, res.fittedvalues), (x, iv_u), (x, iv_l) 다섯 가지 데이터를 각각의 스타일로 그립니다.
  9. legend() 메서드로 범례를 그래프에 배치합니다.
  10. show() 메서드를 호출하여 완성된 그림을 화면에 표시합니다.

예제 코드

import numpy as np
from matplotlib import pyplot as plt
from statsmodels import api as sm
from statsmodels.sandbox.regression.predstd import wls_prediction_std
plt.rcParams["figure.figsize"] = [7.50, 3.50]
plt.rcParams["figure.autolayout"] = True
np.random.seed(9876789)
nsample = 50
sig = 0.5
x = np.linspace(0, 20, nsample)
X = np.column_stack((x, np.sin(x), (x - 5) ** 2, np.ones(nsample)))
beta = [0.5, 0.5, -0.02, 5.]
y_true = np.dot(X, beta)
y = y_true + sig * np.random.normal(size=nsample)
res = sm.OLS(y, X).fit()
prstd, iv_l, iv_u = wls_prediction_std(res)
fig, ax = plt.subplots()
ax.plot(x, y, 'o', label="data")
ax.plot(x, y_true, 'b-', label="True")
ax.plot(x, res.fittedvalues, 'r--.', label="OLS")
ax.plot(x, iv_u, 'r--')
ax.plot(x, iv_l, 'r--')
ax.legend(loc='best')
plt.show()

실행 결과

코드를 실행하면 다음과 같은 그래프가 출력됩니다. 산점도('o' 마커)로 표시된 실제 데이터, 파란색 실선의 참값 곡선(True), 빨간색 점선의 OLS 회귀 피팅 곡선, 그리고 예측 신뢰 구간의 상한과 하한이 한 화면에 함께 나타나므로, 회귀 모델이 데이터를 얼마나 잘 설명하는지 직관적으로 확인할 수 있습니다.

Matplotlib에서 statsmodels 선형 회귀(OLS) 결과를 깔끔하게 시각화하는 방법