ÎncepețiÎncepe gratuit

Validarea încrucișată a unui pipeline pentru modelul de durată a zborurilor

Modelul cu validare încrucișată pe care tocmai l-ai construit era simplu — folosea doar km pentru a prezice duration.

Un alt predictor important al duratei unui zbor este aeroportul de origine. De regulă, zborurile care pleacă din aeroporturi aglomerate durează mai mult până la decolare. Să vedem dacă adăugarea acestui predictor îmbunătățește modelul!

În acest exercițiu vei adăuga câmpul org la model. Deoarece org este o variabilă categorică, sunt necesari câțiva pași suplimentari înainte de a putea fi inclus: trebuie mai întâi transformat într-un index, apoi codificat one-hot, înainte de a fi asamblat împreună cu km pentru a construi modelul de regresie. Vom grupa toate aceste operații într-un pipeline.

Următoarele obiecte au fost deja create:

  • params — un grid de parametri gol
  • evaluator — un evaluator de regresie
  • regression — un obiect LinearRegression cu labelCol='duration'.

Clasele StringIndexer, OneHotEncoder, VectorAssembler și CrossValidator au fost deja importate.

Acest exercițiu face parte din cursul

Machine Learning cu PySpark

Vezi cursul

Instrucțiuni pentru exercițiu

  • Creează un string indexer. Specifică câmpurile de intrare și de ieșire ca org și, respectiv, org_idx.
  • Creează un encoder one-hot. Denumește câmpul de ieșire org_dummy.
  • Asamblează câmpurile km și org_dummy într-un singur câmp numit features.
  • Creează un pipeline folosind următoarele operații: string indexer, encoder one-hot, assembler și regresie liniară. Folosește-l pentru a crea un cross-validator.

Exercițiu interactiv practic

Încearcă acest exercițiu completând acest cod de exemplu.

# Create an indexer for the org field
indexer = ____(____, ____)

# Create an one-hot encoder for the indexed org field
onehot = ____(____, ____)

# Assemble the km and one-hot encoded fields
assembler = ____(____, ____)

# Create a pipeline and cross-validator.
pipeline = ____(stages=[____, ____, ____, ____])
cv = ____(estimator=____,
          estimatorParamMaps=____,
          evaluator=____)
Editează și rulează codul