You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用Apache Beam StatefulDoFn实现EMA时遇TypeError问题求助

问题:Apache Beam StatefulDoFn计算指数移动平均(EMA)时遇到TypeError错误

标签建议

  • apache-beam
  • python
  • stateful-processing
  • exponential-moving-average
  • time-series

问题详情

我正尝试编写一个Apache Beam管道,从REST API导入时序数据,并为每个连续数据点计算指数移动平均(EMA)。最终希望将结果写回REST API,但目前卡在前半部分的计算逻辑上。

从REST API导入数据后,通过FlatMap返回包含时间戳字符串和浮点值的字典,再转换为TimestampedValue传入StatefulDoFn。我原本想通过GlobalWindow保留对之前EMA值的访问,但一直遇到错误:

TypeError: 'float' object is not iterable [while running 'Into windows']

移除DoFn后仅打印输入值是正常的,说明问题出在状态处理逻辑里。以下是我的代码实现:

import argparse
import logging
import os

import apache_beam as beam
from apache_beam.transforms.userstate import BagStateSpec, ReadModifyWriteStateSpec, TimerSpec, on_timer
from apache_beam.transforms.timeutil import TimeDomain
from apache_beam.utils.timestamp import Timestamp
from apache_beam.utils.windowed_value import WindowedValue
from apache_beam.transforms.trigger import AfterCount, AccumulationMode, Repeatedly
from apache_beam.transforms.window import GlobalWindows, Duration, TimestampedValue

import typing
from apache_beam.typehints.decorators import with_input_types, with_output_types
import requests
from apache_beam.options.pipeline_options import PipelineOptions
from typing import Union

# Configure logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

# Load environment variables
django_container_name = os.getenv("DJANGO_CONTAINER", "django")
django_port = os.getenv("DJANGO_PORT", "8000")

def run(argv=None):
    # Parse command-line arguments
    parser = argparse.ArgumentParser()
    args, beam_args = parser.parse_known_args(argv)

    # Define custom pipeline options
    class MyOptions(PipelineOptions):
        @classmethod
        def _add_argparse_args(cls, parser):
            parser.add_argument(
                "--source_table_id",
                required=True,
                help="Table ID containing the raw data.",
                type=str,
            )
            parser.add_argument(
                "--source_table_column",
                required=True,
                help="Column in source table.",
                type=str,
            )
            parser.add_argument(
                "--destination_table_id",
                required=True,
                help="Destination table ID for the output.",
                type=str,
            )
            parser.add_argument(
                "--smoothing",
                required=False,
                help="Smoothing variable of the EMA.",
                type=float,
                default=2.0,
            )
            parser.add_argument(
                "--range",
                required=False,
                help="Timerange of the EMA in minutes.",
                type=int,
                default=60,
            )


    # Parse Beam pipeline options into a PipelineOptions object
    beam_options = PipelineOptions(beam_args)
    args = beam_options.view_as(MyOptions)

    # Function to fetch data from the API
    def get_api_data(dummy_start):

        # Construct the API URL
        api_url = (
            f"http://{django_container_name}:{django_port}"
            f"/api/dynamic-table/{args.source_table_id}/?ordering=timestamp"
        )

        logging.debug("Now fetching from ", api_url)

        # Make the initial API request
        response = requests.get(api_url, timeout=10)
        response = response.json()
        results = response.get("results", [])
        next_url = response.get("next")

        # Fetch data from paginated API
        while next_url:
            logging.debug("Now fetching from ", next_url)
            response = requests.get(next_url, timeout=10)
            response = response.json()
            results.extend(response.get("results", []))
            next_url = response.get("next")

        # Extract relevant data from the API response
        results = [
            {
                "timestamp": result.get("timestamp"),
                args.source_table_column: result.get(args.source_table_column),
            }
            for result in results
        ]

        return results

    def post_api_data(data, feature_name="feature_0"):
        headers = {"Content-Type": "application/json"}
        # Construct the API URL
        api_url = (
            f"http://{django_container_name}:{django_port}"
            f"/api/dynamic-table/feature/{args.destination_table_id}/"
        )

        data[feature_name] = data.pop("ema")
        logging.debug("Now posting to ", api_url)

        # Make the API request
        response = requests.post(api_url, json=data, headers=headers, timeout=10)

        if response.status_code not in [200, 201]:
            logger.error(f"Failed to save data: {response.status_code}, {response.json()}")
            return
        
        return

    class TimestampedValueCoder(beam.coders.Coder):
        def encode(self, value: TimestampedValue):
            """Encode TimestampedValue to bytes."""
            timestamp = value.timestamp.to_rfc3339()
            element = value.value
            return f"{timestamp}:{element}".encode("utf-8")
    
        def decode(self, encoded: bytes):
            decoded = encoded.decode("utf-8")
            timestamp_str, element = decoded.split(":", 1)
            timestamp = Timestamp.from_rfc3339(timestamp_str)
            return TimestampedValue(value=element, timestamp=timestamp)
    
        def is_deterministic(self):
            return True
    
    beam.coders.registry.register_coder(TimestampedValue, TimestampedValueCoder)
    
    class EMAStatefulDoFn(beam.DoFn):

        PREVIOUS_EMA = BagStateSpec("previous_ema", TimestampedValueCoder())

        def process(self, element, 
                    previous_ema=beam.DoFn.StateParam(PREVIOUS_EMA),
                    timestamp=beam.DoFn.TimestampParam,):
            # Get the previous EMA
            previous_ema_value = previous_ema.read()
            if not previous_ema_value:
                previous_ema_value = element

            # Calculate the EMA
            alpha = 2 / (args.smoothing + 1)
            ema = (element * alpha) + (previous_ema_value * (1 - alpha))

            # Update the state
            previous_ema.clear()
            previous_ema.add(ema)

            # Output the EMA
            return [TimestampedValue(value=ema, timestamp=timestamp)]

    # Define the Beam pipeline
    with beam.Pipeline(options=beam_options) as p:

        _ = (
            p
            | "Create" >> beam.Create(["Start"])  # Workaround to kickstart the pipeline
            | "fetch API data" >> beam.FlatMap(get_api_data)  # Fetch data from the API
            | "Timestamp values" >> beam.Map(
                lambda x: TimestampedValue(x[args.source_table_column], 
                                           Timestamp.from_rfc3339(x["timestamp"])))
            | "Into windows" >> beam.WindowInto(GlobalWindows())
            | "Calculate EMA" >> beam.ParDo(EMAStatefulDoFn())
            | "Print EMA" >> beam.Map(logging.info)

        )


if __name__ == "__main__":
    # Set logging level
    logging.getLogger().setLevel(logging.DEBUG)
    run()

问题原因分析

  1. BagStateSpec误用:BagStateSpec用于存储多个值的集合,previous_ema.read()返回的是可迭代对象(而非单个值)。你直接将其与float进行计算,导致迭代错误。
  2. TimestampedValue处理不当:Timestamp values步骤输出的是TimestampedValue对象,你在DoFn中直接用该对象参与数值计算,而非取出其.value属性(实际的浮点值)。
  3. 状态类型选择错误:EMA只需要保存上一个计算值,使用ReadModifyWriteStateSpec(用于单个值的读写更新)比BagStateSpec更合适。

修正后的代码

核心修改部分(EMAStatefulDoFn及相关)

# 移除不必要的TimestampedValueCoder,因为我们将存储浮点值而非TimestampedValue
# beam.coders.registry.register_coder(TimestampedValue, TimestampedValueCoder)

class EMAStatefulDoFn(beam.DoFn):
    # 使用ReadModifyWriteStateSpec存储单个浮点值,用FloatCoder序列化
    PREVIOUS_EMA = ReadModifyWriteStateSpec(
        "previous_ema", beam.coders.FloatCoder())

    def process(self, element, 
                previous_ema=beam.DoFn.StateParam(PREVIOUS_EMA),
                timestamp=beam.DoFn.TimestampParam,):
        # 取出TimestampedValue中的实际浮点值
        current_value = element.value
        # 读取上一个EMA值,初始为None
        last_ema = previous_ema.read()

        if last_ema is None:
            # 第一个数据点,EMA等于当前值
            ema = current_value
        else:
            alpha = 2 / (args.smoothing + 1)
            ema = (current_value * alpha) + (last_ema * (1 - alpha))

        # 更新状态为当前计算的EMA
        previous_ema.write(ema)

        # 输出带时间戳的EMA结果
        return [TimestampedValue(value=ema, timestamp=timestamp)]

其他调整说明

  • 移除了自定义的TimestampedValueCoder,因为现在状态存储的是浮点值,直接用内置的FloatCoder即可。
  • 在DoFn中先提取element.value获取实际数值,避免对象直接参与计算。
  • 用ReadModifyWriteStateSpec替代BagStateSpec,简化单个值的状态管理逻辑。

窗口选择建议

如果你的需求是持续累积的全局EMA(基于所有历史数据点),那么GlobalWindow+状态的方案是合理的;如果需要基于滑动时间窗口的EMA(仅计算最近N分钟内数据的EMA),则可以改用SlidingWindows,同时结合状态或窗口聚合逻辑,但需要注意窗口触发和清理的规则。

内容的提问来源于stack exchange,提问作者jrood

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.14 01:38:09