Skip to content

Commit 2747573

Browse files
timsaucerclaude
andcommitted
feat: add DataFrame.fill_nan
Replace NaN in floating-point columns, optionally limited to a subset. Mirrors fill_null and wraps upstream DataFrame::fill_nan. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 parent e8c2192 commit 2747573

3 files changed

Lines changed: 92 additions & 0 deletions

File tree

‎crates/core/src/dataframe.rs‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1320,6 +1320,26 @@ impl PyDataFrame {
13201320
let df = self.df.as_ref().fill_null(&scalar_value.0, &cols)?;
13211321
Ok(Self::new(df))
13221322
}
1323+
1324+
/// Fill NaN values with a specified value for specific floating-point columns
1325+
#[pyo3(signature = (value, columns=None))]
1326+
fn fill_nan(
1327+
&self,
1328+
value: Py<PyAny>,
1329+
columns: Option<Vec<PyBackedStr>>,
1330+
py: Python,
1331+
) -> PyDataFusionResult<Self> {
1332+
let scalar_value: PyScalarValue = value.extract(py)?;
1333+
1334+
let cols = match columns {
1335+
Some(col_names) => col_names.iter().map(|c| c.to_string()).collect(),
1336+
None => Vec::new(), // Empty vector means fill NaN for all columns
1337+
};
1338+
1339+
let cols = cols.iter().map(String::as_str).collect::<Vec<_>>();
1340+
let df = self.df.as_ref().fill_nan(&scalar_value.0, &cols)?;
1341+
Ok(Self::new(df))
1342+
}
13231343
}
13241344

13251345
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]

‎python/datafusion/dataframe.py‎

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1875,6 +1875,34 @@ def fill_null(self, value: Any, subset: list[str] | None = None) -> DataFrame:
18751875
"""
18761876
return DataFrame(self.df.fill_null(value, subset))
18771877

1878+
def fill_nan(self, value: float, subset: list[str] | None = None) -> DataFrame:
1879+
"""Fill NaN values in floating-point columns with a value.
1880+
1881+
Only floating-point columns are changed; others are kept unchanged, as is
1882+
any column ``value`` cannot be cast to. NaN is distinct from null, which
1883+
:py:meth:`fill_null` handles.
1884+
1885+
Args:
1886+
value: Value to replace NaN with. Will be cast to match column type.
1887+
subset: Optional list of column names to fill. If None, fills all
1888+
floating-point columns.
1889+
1890+
Returns:
1891+
DataFrame with NaN values replaced.
1892+
1893+
Examples:
1894+
>>> from datafusion import SessionContext
1895+
>>> ctx = SessionContext()
1896+
>>> nan = float("nan")
1897+
>>> df = ctx.from_pydict({"a": [1.0, nan, None], "b": [nan, 2.0, 3.0]})
1898+
>>> df.fill_nan(0.0).to_pydict()
1899+
{'a': [1.0, 0.0, None], 'b': [0.0, 2.0, 3.0]}
1900+
1901+
>>> df.fill_nan(0.0, subset=["a"]).collect_column("b")[0].as_py()
1902+
nan
1903+
"""
1904+
return DataFrame(self.df.fill_nan(value, subset))
1905+
18781906

18791907
class InsertOp(Enum):
18801908
"""Insert operation mode.

‎python/tests/test_dataframe.py‎

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3450,6 +3450,50 @@ def test_fill_null_all_null_column(ctx):
34503450
assert result.column(1).to_pylist() == ["filled", "filled", "filled"]
34513451

34523452

3453+
def _nan_df(ctx):
3454+
nan = float("nan")
3455+
batch = pa.RecordBatch.from_arrays(
3456+
[
3457+
pa.array([1.0, nan, None], type=pa.float64()),
3458+
pa.array([nan, 2.0, 3.0], type=pa.float32()),
3459+
pa.array([1, 2, 3]),
3460+
pa.array(["x", "nan", None]),
3461+
],
3462+
names=["f64", "f32", "i", "s"],
3463+
)
3464+
return ctx.create_dataframe([[batch]])
3465+
3466+
3467+
def _is_nan(v):
3468+
return v is not None and v != v # noqa: PLR0124
3469+
3470+
3471+
def test_fill_nan_all_columns(ctx):
3472+
result = _nan_df(ctx).fill_nan(0.0).to_pydict()
3473+
# NaN replaced in both float widths; null is not NaN and stays null.
3474+
assert result["f64"] == [1.0, 0.0, None]
3475+
assert result["f32"] == [0.0, 2.0, 3.0]
3476+
# Non-float columns are untouched.
3477+
assert result["i"] == [1, 2, 3]
3478+
assert result["s"] == ["x", "nan", None]
3479+
3480+
3481+
def test_fill_nan_subset(ctx):
3482+
result = _nan_df(ctx).fill_nan(-1.0, subset=["f32"]).to_pydict()
3483+
assert result["f32"] == [-1.0, 2.0, 3.0]
3484+
assert _is_nan(result["f64"][1])
3485+
3486+
3487+
def test_fill_nan_preserves_schema(ctx):
3488+
df = _nan_df(ctx)
3489+
assert df.fill_nan(0.0).schema() == df.schema()
3490+
3491+
3492+
def test_fill_nan_unknown_column_raises(ctx):
3493+
with pytest.raises(Exception, match="missing"):
3494+
_nan_df(ctx).fill_nan(0.0, subset=["missing"]).collect()
3495+
3496+
34533497
_slow_udf_started = threading.Event()
34543498

34553499

0 commit comments

Comments
 (0)