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

Matplotlib로 그린 그래프 코드의 단위 테스트 작성 방법

코드에 대한 단위 테스트 케이스를 작성할 때는 배열을 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() 메서드를 활용하면, 시각화 결과물을 눈으로 확인하지 않고도 프로그래밍 방식으로 그래프 데이터의 정확성을 자동으로 검증할 수 있습니다. 이는 데이터 시각화 파이프라인의 신뢰성을 확보하는 데 매우 유용한 접근 방식입니다.