mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-08-03 18:27:43 +00:00
fix(auto-fixer): add five new factor code fixes for groupby/apply errors
1. groupby(level=['instrument','date']) → get_level_values() — string level names like 'date' don't exist in the (datetime, instrument) MultiIndex; replaced with get_level_values(0).normalize() + get_level_values(1). 2. groupby(level=['date','instrument']) — symmetric fix for reversed order. 3. groupby(level=['instrument']) → groupby(level=1) — single string level. 4. groupby(level=N)['col'].apply(lambda) → transform(lambda) — apply() on a grouped Series prepends an extra index level, causing index shape mismatch when assigned back; transform() preserves the original index. 5. df.loc[instrument] DateParseError fix (instrument_loc_multiindex) — already committed, adding supporting tests for groupby(level=['instrument','date']). Adds 5 new tests (TestGroupbyLevelStringNames, TestGroupbyApplyToTransform) — total 28 tests, all passing. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -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()")
|
||||
|
||||
Reference in New Issue
Block a user