코드에 대한 단위 테스트 케이스를 작성할 때는 배열을 x축 데이터로 받아 y = x² 곡선을 그리는 플롯을 예시로 들 수 있습니다. 테스트 과정에서는 x 데이터 포인트에 해당하는 y_data를 추출하여 검증하게 됩니다.
테스트 절차
- x 값과 x² 값을 plot() 메서드로 그리고 플롯을 반환하는 메서드, 즉 plot_sqr_curve(x)를 생성합니다.
- 테스트에는 unittest.TestCase를 활용합니다.
- 다음 내용을 포함하는 test_curve_sqr_plot() 메서드를 작성합니다.
- 곡선을 그리기 위한 x 데이터 포인트를 생성합니다.
- 위에서 만든 x 데이터 포인트를 이용해 y 데이터 포인트를 생성합니다.
- x와 y 데이터 포인트로 곡선을 플롯합니다.
- 플롯 객체(pt)에서 x와 y 데이터를 추출합니다.
- 주어진 표현식이 참인지 여부를 확인합니다.
예제 코드
import unittest
import numpy as np
from matplotlib import pyplot as plt
def plot_sqr_curve(x):
"""
y = x^2 곡선을 그리는 함수.
"""
return plt.plot(x, np.square(x))
class TestSqrCurve(unittest.TestCase):
def test_curve_sqr_plot(self):
x = np.array([1, 3, 4])
y = np.square(x)
pt, = plot_sqr_curve(x)
y_data = pt.get_data()[1]
x_data = pt.get_data()[0]
self.assertTrue((y == y_data).all())
self.assertTrue((x == x_data).all())
if __name__ == '__main__':
unittest.main()실행 결과
Ran 1 test in 1.587s OK
코드 설명
이 테스트 코드의 핵심 동작 원리는 다음과 같습니다.
- plot_sqr_curve(x): 입력받은 x 배열을 제곱하여 y값으로 사용하고, matplotlib의
plt.plot()으로 곡선을 그린 뒤 라인 객체를 반환합니다. - pt.get_data(): 반환된 라인 객체에서 실제로 그려진 x, y 좌표 데이터를 튜플 형태로 가져옵니다.
- assertTrue(): 계산된 y값과 실제 플롯에 반영된 y_data가 일치하는지, 그리고 x값도 올바르게 전달되었는지 검증합니다.
이처럼 matplotlib의 Line2D 객체가 제공하는 get_data() 메서드를 활용하면, 시각화 결과물을 눈으로 확인하지 않고도 프로그래밍 방식으로 그래프 데이터의 정확성을 자동으로 검증할 수 있습니다. 이는 데이터 시각화 파이프라인의 신뢰성을 확보하는 데 매우 유용한 접근 방식입니다.