始める無料で始める

数値フィールドを条件付きでフィルタリングする

データの文脈を理解することは非常に重要です。ここでは、住宅が通常どの価格帯で売れているかを把握したいと考えています。平均から大きく外れた価格で売れた外れ値の物件は除外しましょう。平均と標準偏差を計算し、それらを使ってほぼ正規分布に近いフィールド log_SalesClosePrice をフィルタします。

この演習はコースの一部です

PySparkで学ぶ特徴量エンジニアリング

コースを見る

演習の手順

  • pyspark.sql.functions から mean()stddev() をインポートします。
  • インポートした関数を使って、agg()'log_SalesClosePrice' の平均と標準偏差を計算します。
  • mean_val ± stddev_val の3倍で上下限を作成します。
  • low_boundhi_bound の両方を使って、'log_SalesClosePrice' に対する where() フィルタを作成します。

実践的なインタラクティブ演習

このサンプルコードを完成させて、この演習に挑戦してみましょう。

from ____ import ____, ____

# Calculate values used for outlier filtering
mean_val = df.____({____: ____}).collect()[0][0]
stddev_val = df.____({____: ____}).collect()[0][0]

# Create three standard deviation (μ ± 3σ) lower and upper bounds for data
low_bound = ____ - (3 * ____)
hi_bound = ____ + (3 * ____)

# Filter the data to fit between the lower and upper bounds
df = df.____((df[____] < ____) ____ (df[____] > ____))
コードを編集して実行