diff --git a/rdagent/components/coder/factor_coder/auto_fixer.py b/rdagent/components/coder/factor_coder/auto_fixer.py index 6c7b1b78..b9a277d5 100644 --- a/rdagent/components/coder/factor_coder/auto_fixer.py +++ b/rdagent/components/coder/factor_coder/auto_fixer.py @@ -654,6 +654,25 @@ class FactorAutoFixer: f"groupby: {df_var}.groupby(level={level})['{col}'].apply() → transform()" ) + # === FIX: .transform(...).reset_index(level=N, drop=True) === + # transform() already returns the same index as the input — adding reset_index() + # after it drops an index level and causes ValueError on assignment back to df['col']. + # Detected line-by-line: if a line contains both .transform( and .reset_index(level= + reset_suffix = re.compile(r'\s*\.reset_index\s*\(\s*level\s*=[^,)]+,\s*drop\s*=\s*True\s*\)\s*$') + new_lines = [] + changed = False + for line in fixed_code.splitlines(): + if '.transform(' in line and '.reset_index(' in line: + cleaned = reset_suffix.sub('', line) + if cleaned != line: + new_lines.append(cleaned) + changed = True + continue + new_lines.append(line) + if changed: + fixed_code = '\n'.join(new_lines) + self.fixes_applied.append("groupby: removed spurious .reset_index() after .transform()") + # Pattern: Simple groupby().apply() with rolling().method() # df.groupby(level=N).apply(lambda x: x['col'].rolling(...).method()) apply_pattern = r"df\.groupby\(level=(\d+)\)\.apply\(\s*lambda\s+x:\s+x\['([^']+)'\]\.rolling\([^)]+\)\.(\w+)\([^)]*\)\s*\)" diff --git a/test/qlib/test_auto_fixer.py b/test/qlib/test_auto_fixer.py index e56c27d6..09698ef6 100644 --- a/test/qlib/test_auto_fixer.py +++ b/test/qlib/test_auto_fixer.py @@ -197,6 +197,13 @@ class TestGroupbyApplyToTransform: assert "lambda x: x.cumsum()" in result assert ".transform(" in result + def test_transform_reset_index_stripped(self, fixer): + # .transform() already preserves index — .reset_index() after it is wrong + code = "df['v'] = df.groupby(level=1)['x'].transform(lambda x: x.rolling(20).mean()).reset_index(level=0, drop=True)" + result = fixer.fix(code) + assert ".reset_index(level=0, drop=True)" not in result + assert ".transform(" in result + class TestRollingDdof: def test_removes_ddof_from_rolling_args(self, fixer):