For example in Pandas I would do
data_df = (
pd.DataFrame(dict(col1=['a', 'b', 'c'], col2=['1', '2', '3']))
.pipe(lambda df: df[df.col1 != 'a'])
)
This is similar to R's pipe %>%
Is there something similar in PySpark?
For example in Pandas I would do
data_df = (
pd.DataFrame(dict(col1=['a', 'b', 'c'], col2=['1', '2', '3']))
.pipe(lambda df: df[df.col1 != 'a'])
)
This is similar to R's pipe %>%
Is there something similar in PySpark?
You can define a "pandas-like" pipe method and bind it to the DataFrame class:
from pyspark.sql import DataFrame
def pipe(self, func, *args, **kwargs):
return func(self, *args, **kwargs)
DataFrame.pipe = pipe
Then, you can pass functions to the pipe method to apply to the pyspark DataFrame. For instance, suppose that you want to select all the columns from a DataFrame my_df, except for the last two, after having changed its columns. You can use pipe for this:
my_new_df = (
my_df
# Perform some operations to add and/or remove columns
...
# At this point the list of columns is different
# from `my_df.columns`
.pipe(lambda df: df.select(*df.columns[:-2]))
)
In PySpark the pipe function is called transform with documentation here
The behavior is identical to the Pandas pipe operator.
So the example in PySpark would look like
data_df = (
spark.createDataFrame(pd.DataFrame(dict(col1=['a', 'b', 'c'], col2=['1', '2', '3'])))
.transform(lambda df: df.filter("col1 != 'a'"))
)
I think, in pyspark, you can easily achieve this pipe functionality with help of pipeline.
Example: Let's take the example you provided
val df = Seq(("a", 1), ("b", 2), ("c", 3)).toDF("col1", "col2")
df.show(false)
df.printSchema()
/**
* +----+----+
* |col1|col2|
* +----+----+
* |a |1 |
* |b |2 |
* |c |3 |
* +----+----+
*
* root
* |-- col1: string (nullable = true)
* |-- col2: integer (nullable = false)
*/
for .pipe(lambda df: df[df.col1 != 'a']), we can easily use spark SQLTransformer. so no need to create custom transformer
val transform1 = new SQLTransformer()
.setStatement("select * from __THIS__ where col1 != 'a'")
val transform2 = new SQLTransformer()
.setStatement("select col1, col2, SQRT(col2) as col3 from __THIS__")
val pipeline = new Pipeline()
.setStages(Array(transform1, transform2))
pipeline.fit(df).transform(df)
.show(false)
/**
* +----+----+------------------+
* |col1|col2|col3 |
* +----+----+------------------+
* |b |2 |1.4142135623730951|
* |c |3 |1.7320508075688772|
* +----+----+------------------+
*/