@@ -16,7 +16,7 @@
1616# under the License.
17171818importdatafusion
19-importpyarrow
19+importpyarrowaspa
2020importpyarrow.compute
2121fromdatafusionimportAccumulator, col, udaf
2222@@ -26,48 +26,44 @@ class MyAccumulator(Accumulator):
2626 Interface of a user-defined accumulation.
2727 """
282829-def__init__(self):
30-self._sum=pyarrow.scalar(0.0)
29+def__init__(self)->None:
30+self._sum=pa.scalar(0.0)
313132-defupdate(self, values: pyarrow.Array) ->None:
32+defupdate(self, values: pa.Array) ->None:
3333# not nice since pyarrow scalars can't be summed yet. This breaks on `None`
34-self._sum=pyarrow.scalar(
35-self._sum.as_py() +pyarrow.compute.sum(values).as_py()
36- )
34+self._sum=pa.scalar(self._sum.as_py() +pa.compute.sum(values).as_py())
373538-defmerge(self, states: pyarrow.Array) ->None:
36+defmerge(self, states: pa.Array) ->None:
3937# not nice since pyarrow scalars can't be summed yet. This breaks on `None`
40-self._sum=pyarrow.scalar(
41-self._sum.as_py() +pyarrow.compute.sum(states).as_py()
42- )
38+self._sum=pa.scalar(self._sum.as_py() +pa.compute.sum(states).as_py())
433944-defstate(self) ->pyarrow.Array:
45-returnpyarrow.array([self._sum.as_py()])
40+defstate(self) ->pa.Array:
41+returnpa.array([self._sum.as_py()])
464247-defevaluate(self) ->pyarrow.Scalar:
43+defevaluate(self) ->pa.Scalar:
4844returnself._sum
494550465147# create a context
5248ctx=datafusion.SessionContext()
53495450# create a RecordBatch and a new DataFrame from it
55-batch=pyarrow.RecordBatch.from_arrays(
56- [pyarrow.array([1, 2, 3]), pyarrow.array([4, 5, 6])],
51+batch=pa.RecordBatch.from_arrays(
52+ [pa.array([1, 2, 3]), pa.array([4, 5, 6])],
5753names=["a", "b"],
5854)
5955df=ctx.create_dataframe([[batch]])
60566157my_udaf=udaf(
6258MyAccumulator,
63-pyarrow.float64(),
64-pyarrow.float64(),
65- [pyarrow.float64()],
59+pa.float64(),
60+pa.float64(),
61+ [pa.float64()],
6662"stable",
6763)
68646965df=df.aggregate([], [my_udaf(col("a"))])
70667167result=df.collect()[0]
726873-assertresult.column(0) ==pyarrow.array([6.0])
69+assertresult.column(0) ==pa.array([6.0])