diff --git a/rdagent/components/coder/factor_coder/auto_fixer.py b/rdagent/components/coder/factor_coder/auto_fixer.py index 7aec60d9..35e2afaa 100644 --- a/rdagent/components/coder/factor_coder/auto_fixer.py +++ b/rdagent/components/coder/factor_coder/auto_fixer.py @@ -258,6 +258,39 @@ class FactorAutoFixer: fixed_code = re.sub(r"\.groupby\(\['instrument'\]\)", ".groupby(level=1)", fixed_code) self.fixes_applied.append("multiindex_groupby: groupby(['instrument']) → groupby(level=1)") + # groupby(level=['instrument', 'date']) — uses level= keyword with string names. + # 'date' is NOT a valid level name in our (datetime, instrument) MultiIndex; + # replace with get_level_values to normalize datetime to daily timestamps. + fixed_code = re.sub( + r"(\w+)\.groupby\(level=\['instrument',\s*'date'\]\)", + lambda m: ( + self.fixes_applied.append( + f"multiindex_groupby: {m.group(0)[:60]} → two-level get_level_values" + ) + or f"{m.group(1)}.groupby([{m.group(1)}.index.get_level_values(1), " + f"{m.group(1)}.index.get_level_values(0).normalize()])" + ), + fixed_code, + ) + # groupby(level=['date', 'instrument']) + fixed_code = re.sub( + r"(\w+)\.groupby\(level=\['date',\s*'instrument'\]\)", + lambda m: ( + self.fixes_applied.append( + f"multiindex_groupby: {m.group(0)[:60]} → two-level get_level_values" + ) + or f"{m.group(1)}.groupby([{m.group(1)}.index.get_level_values(0).normalize(), " + f"{m.group(1)}.index.get_level_values(1)])" + ), + fixed_code, + ) + # single: groupby(level=['instrument']) → groupby(level=1) + fixed_code = re.sub( + r"\.groupby\(level=\['instrument'\]\)", + lambda m: (self.fixes_applied.append("multiindex_groupby: groupby(level=['instrument']) → level=1") or ".groupby(level=1)"), + fixed_code, + ) + return fixed_code def _fix_chained_groupby(self, code: str) -> str: @@ -584,6 +617,26 @@ class FactorAutoFixer: fixed_code = fixed_code.replace(old_code, new_code) self.fixes_applied.append(f"groupby: fixed rolling correlation (window={window}) with reset_index") + # === GENERAL FIX: DF.groupby(level=N)['col'].apply(lambda x: EXPR) === + # apply() on a grouped Series returns a MultiIndex result (extra level prepended), + # causing index shape mismatch when assigned back to df['col']. + # Replace with transform() which preserves the original index. + col_apply_pattern = re.compile( + r"(\w+)\.groupby\(level=(\d+)\)\['([^']+)'\]\.apply\((\s*lambda\s+\w+\s*:.*?)\)", + re.DOTALL, + ) + for m in list(col_apply_pattern.finditer(fixed_code)): + full = m.group(0) + df_var = m.group(1) + level = m.group(2) + col = m.group(3) + lam = m.group(4).strip() + new_expr = f"{df_var}.groupby(level={level})['{col}'].transform({lam})" + fixed_code = fixed_code.replace(full, new_expr, 1) + self.fixes_applied.append( + f"groupby: {df_var}.groupby(level={level})['{col}'].apply() → 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 4db980ef..5561da81 100644 --- a/test/qlib/test_auto_fixer.py +++ b/test/qlib/test_auto_fixer.py @@ -156,6 +156,41 @@ class TestInstrumentLocMultiindex: assert "df.loc[date]" in result +class TestGroupbyLevelStringNames: + def test_level_instrument_date_replaced(self, fixer): + code = "df.groupby(level=['instrument', 'date'])['col'].transform('sum')" + result = fixer.fix(code) + assert "get_level_values(1)" in result + assert "get_level_values(0).normalize()" in result + assert "level=['instrument', 'date']" not in result + + def test_level_date_instrument_replaced(self, fixer): + code = "data.groupby(level=['date', 'instrument'])['x'].mean()" + result = fixer.fix(code) + assert "get_level_values(0).normalize()" in result + assert "get_level_values(1)" in result + + def test_level_instrument_single_replaced(self, fixer): + code = "df.groupby(level=['instrument'])['vol'].sum()" + result = fixer.fix(code) + assert "groupby(level=1)" in result + assert "level=['instrument']" not in result + + +class TestGroupbyApplyToTransform: + def test_col_apply_lambda_replaced(self, fixer): + code = "df_overlap.groupby(level=1)['$close'].apply(lambda x: np.log(x / x.shift(1)))" + result = fixer.fix(code) + assert ".transform(" in result + assert ".apply(" not in result + + def test_col_apply_lambda_preserves_lambda_body(self, fixer): + code = "series.groupby(level=1)['ret'].apply(lambda x: x.cumsum())" + result = fixer.fix(code) + assert "lambda x: x.cumsum()" in result + assert ".transform(" in result + + class TestRollingDdof: def test_removes_ddof_from_rolling_args(self, fixer): result = fixer.fix("df.rolling(20, min_periods=1, ddof=1).std()")