mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-29 16:37:43 +00:00
Compare commits
71 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b53749df7d | |||
| 3a1a3d5f77 | |||
| 13cbd42ecf | |||
| 64ed6b0cce | |||
| bf36f54159 | |||
| 79f1d34083 | |||
| 150a818e07 | |||
| b6d1caecc9 | |||
| 73e600bf25 | |||
| 9960633d01 | |||
| 3522a2eca1 | |||
| a5f091f1ca | |||
| 528d470754 | |||
| 910fbea27e | |||
| ab3f5f111d | |||
| a910d70d40 | |||
| 31a75eeb07 | |||
| 11f5dadd2d | |||
| 51a624c31e | |||
| 9947ea3928 | |||
| bc96d26371 | |||
| c6e8f3d3a3 | |||
| 35d2b81158 | |||
| 4fd5117af6 | |||
| 96d6923433 | |||
| ef12b33aca | |||
| a1e9417658 | |||
| a65ab828c4 | |||
| 840e12e6aa | |||
| 1d1b7b6984 | |||
| 4f1660b6aa | |||
| a52adf5b5a | |||
| 537f730c93 | |||
| 9c07b07995 | |||
| 9a47691420 | |||
| a370690ee8 | |||
| eaebd60d93 | |||
| 8070de3ae1 | |||
| ce806ea60b | |||
| 57e2609402 | |||
| bb32276332 | |||
| 35d03a8a0d | |||
| 9591c11702 | |||
| 7582e55bb3 | |||
| 27803e8b85 | |||
| 944af06a87 | |||
| 97e42d7a1a | |||
| 5481e83f03 | |||
| 88c4cc4a33 | |||
| 01889a6b64 | |||
| 443c6d47b2 | |||
| b10d3512df | |||
| d75cba934e | |||
| 38fa760429 | |||
| 0ce6f6ec6d | |||
| d17d424ee9 | |||
| 5d8e53d208 | |||
| 7880a9315a | |||
| 17bba1a920 | |||
| 360df4083a | |||
| e2c2fefe9a | |||
| 32f7d66e07 | |||
| 4d6ef04411 | |||
| a757eb4b79 | |||
| 1dc44d33ed | |||
| e672433305 | |||
| f875439105 | |||
| 4afd03dbaf | |||
| 90b36c174d | |||
| 5da3ba6752 | |||
| d8b5bc4237 |
@@ -14,7 +14,7 @@ jobs:
|
||||
security:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- name: Run Bandit (Security Scan)
|
||||
uses: PyCQA/bandit-action@v1
|
||||
@@ -25,9 +25,9 @@ jobs:
|
||||
test:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
- uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: "3.10"
|
||||
cache: "pip"
|
||||
|
||||
@@ -36,11 +36,11 @@ jobs:
|
||||
steps:
|
||||
# Checkout the repository to the GitHub Actions runner
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v6
|
||||
|
||||
# Execute Codacy Analysis CLI and generate a SARIF output with the security issues identified during the analysis
|
||||
- name: Run Codacy Analysis CLI
|
||||
uses: codacy/codacy-analysis-cli-action@d840f886c4bd4edc059706d09c6a1586111c540b
|
||||
uses: codacy/codacy-analysis-cli-action@562ee3e92b8e92df8b67e0a5ff8aa8e261919c08
|
||||
env:
|
||||
JAVA_TOOL_OPTIONS: "-Dfile.encoding=UTF-8"
|
||||
with:
|
||||
|
||||
@@ -46,7 +46,7 @@ jobs:
|
||||
name: Validate Commit Messages
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
|
||||
@@ -25,10 +25,10 @@ jobs:
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v6
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
@@ -64,7 +64,7 @@ jobs:
|
||||
|
||||
- name: Upload docs artifact
|
||||
if: github.ref == 'refs/heads/main'
|
||||
uses: actions/upload-pages-artifact@v3
|
||||
uses: actions/upload-pages-artifact@v5
|
||||
with:
|
||||
path: docs/_build/html
|
||||
|
||||
|
||||
@@ -16,10 +16,10 @@ jobs:
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v6
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
|
||||
@@ -12,7 +12,8 @@ jobs:
|
||||
release-please:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: googleapis/release-please-action@v4
|
||||
- uses: googleapis/release-please-action@v5
|
||||
with:
|
||||
release-type: python
|
||||
token: ${{ secrets.GITHUB_TOKEN }}
|
||||
config-file: release-please-config.json
|
||||
manifest-file: .release-please-manifest.json
|
||||
|
||||
@@ -19,9 +19,9 @@ jobs:
|
||||
python-version: ["3.10", "3.11"]
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
- uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
cache: "pip"
|
||||
@@ -49,9 +49,9 @@ jobs:
|
||||
name: Dependency Audit
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
- uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: "3.10"
|
||||
cache: "pip"
|
||||
|
||||
@@ -19,10 +19,10 @@ jobs:
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v6
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
{
|
||||
".": "1.3.10"
|
||||
}
|
||||
+132
@@ -1,5 +1,137 @@
|
||||
# Changelog
|
||||
|
||||
## [1.3.10](https://github.com/TPTBusiness/Predix/compare/v1.3.9...v1.3.10) (2026-05-01)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **security:** replace remaining assert statements with proper error handling ([928533d](https://github.com/TPTBusiness/Predix/commit/928533d9a81bd5062f07458fbf94d3c7fe347775))
|
||||
|
||||
## [1.3.9](https://github.com/TPTBusiness/Predix/compare/v1.3.8...v1.3.9) (2026-05-01)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **security:** resolve path-injection, B701, B101, B112 Bandit alerts ([20b89a0](https://github.com/TPTBusiness/Predix/commit/20b89a061843b39836e975f158404e8e2d4627cd))
|
||||
|
||||
## [1.3.8](https://github.com/TPTBusiness/Predix/compare/v1.3.7...v1.3.8) (2026-04-30)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **deps:** relax aiohttp constraint to >=3.13.4 for litellm compatibility ([34ab192](https://github.com/TPTBusiness/Predix/commit/34ab1923a887089eb36e5cbad6cb8df16f0333ca))
|
||||
* **qlib:** correct indentation in except blocks in quant_proposal and factor_runner ([8143451](https://github.com/TPTBusiness/Predix/commit/8143451e8c0ead01c4d86d19669268c7bfb15fac))
|
||||
* **security:** replace eval() with ast.literal_eval in finetune validator (B307) ([0508caf](https://github.com/TPTBusiness/Predix/commit/0508caf9140d210b823fefefa28ee535ec85a0ae))
|
||||
* **security:** replace shell=True subprocess calls with list args in env.py (B602) ([2012d5a](https://github.com/TPTBusiness/Predix/commit/2012d5ae4e77cc2f1ab9a48beaaac5a74695d083))
|
||||
* **security:** resolve path-injection and add nosec for safe temp paths (B108, py/path-injection) ([6727480](https://github.com/TPTBusiness/Predix/commit/67274803bd1d14e5d1df9a063f46b2edb8501a2b))
|
||||
|
||||
## [1.3.7](https://github.com/TPTBusiness/Predix/compare/v1.3.6...v1.3.7) (2026-04-30)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **security:** nosec for B608/B701 false positives in UI and template code ([5eb5d7e](https://github.com/TPTBusiness/Predix/commit/5eb5d7e8fdbe90e0dced83fef4e09f5a33e96b2b))
|
||||
* **security:** replace eval() with ast.literal_eval and add request timeouts (B307, B113) ([3301ada](https://github.com/TPTBusiness/Predix/commit/3301ada697ca7d3afa1a188d2a76a87ae98b4529))
|
||||
* **security:** replace shell=True subprocess calls with list args (B602) ([13c08f4](https://github.com/TPTBusiness/Predix/commit/13c08f4ce6813eb7c314087921ec8c0f40074bd7))
|
||||
|
||||
## [1.3.6](https://github.com/TPTBusiness/Predix/compare/v1.3.5...v1.3.6) (2026-04-30)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **security:** real fix for B110 (logging in factor_proposal.py [#746](https://github.com/TPTBusiness/Predix/issues/746)) ([16624e0](https://github.com/TPTBusiness/Predix/commit/16624e0bd966ae4d24c4a3eb42bbc31c11da3136))
|
||||
* **security:** real fix for B110 (logging in factor_runner.py [#744](https://github.com/TPTBusiness/Predix/issues/744)) ([88cf0fb](https://github.com/TPTBusiness/Predix/commit/88cf0fb8828b11c97f2f3ae2881a4900b020c6f0))
|
||||
* **security:** real fix for B110 (logging in quant_proposal.py [#741](https://github.com/TPTBusiness/Predix/issues/741)) ([7cf2a64](https://github.com/TPTBusiness/Predix/commit/7cf2a644f553b054bd4b0607ea51e5372e68d90a))
|
||||
* **security:** real fix for B110 (logging in quant_proposal.py [#741](https://github.com/TPTBusiness/Predix/issues/741)) ([ef985f8](https://github.com/TPTBusiness/Predix/commit/ef985f86035d8dca707c60137e6508349a0c4ae6))
|
||||
* **security:** real fix for B404/B603 (sys.executable in factor_runner.py [#745](https://github.com/TPTBusiness/Predix/issues/745)) ([819655a](https://github.com/TPTBusiness/Predix/commit/819655aaa3efa76596d60501d0e8ca365df3e5e2))
|
||||
* **security:** revert broken read_pickle encoding arg in kaggle template (B301) ([3574907](https://github.com/TPTBusiness/Predix/commit/35749073c91e69f63ddaad61dae3f2b799327e63))
|
||||
* **security:** validate SQL identifiers in _add_column_if_not_exists (B608) ([e10dfa2](https://github.com/TPTBusiness/Predix/commit/e10dfa2576038e911f83595d3b466c261bc0cd54))
|
||||
* **security:** whitelist-validate metric column in get_top_factors (B608) ([e50519f](https://github.com/TPTBusiness/Predix/commit/e50519fe066e68aec2f19b83df4f643c3c22053d))
|
||||
|
||||
## [1.3.5](https://github.com/TPTBusiness/Predix/compare/v1.3.4...v1.3.5) (2026-04-27)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **auto-fixer:** add five new factor code fixes for groupby/apply errors ([449c8fd](https://github.com/TPTBusiness/Predix/commit/449c8fd70a327e604dcca122e4a134f0cca918e4))
|
||||
* **auto-fixer:** add four new factor code fixes for common runtime errors ([40484f6](https://github.com/TPTBusiness/Predix/commit/40484f6d300425da481f1edd325da4acbc06ec7d))
|
||||
* **auto-fixer:** add groupby([level=N,'date']) SyntaxError fix ([ca77c00](https://github.com/TPTBusiness/Predix/commit/ca77c005bea4abdd8854c1de2b0e8d03b7742161))
|
||||
* **auto-fixer:** disable _fix_min_periods for intraday data ([77b0740](https://github.com/TPTBusiness/Predix/commit/77b0740f059349df7e769a378af728aa33b2070e))
|
||||
* **auto-fixer:** fix chained groupby(level=N).groupby('date') pattern ([7d5fe32](https://github.com/TPTBusiness/Predix/commit/7d5fe32b31a19ce8b04bd8f5a430720fdb748f7a))
|
||||
* **auto-fixer:** fix df.loc[instrument] DateParseError on MultiIndex frames ([b7860ea](https://github.com/TPTBusiness/Predix/commit/b7860eafc0ad26384947ce0510ecf4e9f3425807))
|
||||
* **auto-fixer:** fix df['instrument'] KeyError on MultiIndex frames ([aad6bd1](https://github.com/TPTBusiness/Predix/commit/aad6bd1c7c720b3d486e0cf248337f32394773b1))
|
||||
* **auto-fixer:** fix two assignment-target bugs in instrument column fixers ([421eedf](https://github.com/TPTBusiness/Predix/commit/421eedffed4b883c24397dc5581c019a3985277f))
|
||||
* **auto-fixer:** preserve date dimension in groupby(['instrument','date']) fix ([b58fdd8](https://github.com/TPTBusiness/Predix/commit/b58fdd8be43720b5d4363e0f8de9a01591d4d2dc))
|
||||
* **auto-fixer:** remove ddof from rolling() args, not only from std()/var() ([b0fc328](https://github.com/TPTBusiness/Predix/commit/b0fc328d0d4a041c65d8eeb32cb3f2bb86568406))
|
||||
* **auto-fixer:** strip spurious .reset_index() after .transform() calls ([8708aae](https://github.com/TPTBusiness/Predix/commit/8708aae6e08728cda1875c775a76dc92e43576f3))
|
||||
* **loop:** prevent step_idx advance on unhandled exceptions + fix consecutive assistant messages ([5ec4ad1](https://github.com/TPTBusiness/Predix/commit/5ec4ad1b96b5b99ef42bea7bb828cb1ef709a688))
|
||||
|
||||
## [1.3.4](https://github.com/TPTBusiness/Predix/compare/v1.3.3...v1.3.4) (2026-04-27)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **auto-fixer:** add five new factor code fixes for groupby/apply errors ([449c8fd](https://github.com/TPTBusiness/Predix/commit/449c8fd70a327e604dcca122e4a134f0cca918e4))
|
||||
* **auto-fixer:** add four new factor code fixes for common runtime errors ([40484f6](https://github.com/TPTBusiness/Predix/commit/40484f6d300425da481f1edd325da4acbc06ec7d))
|
||||
* **auto-fixer:** add groupby([level=N,'date']) SyntaxError fix ([ca77c00](https://github.com/TPTBusiness/Predix/commit/ca77c005bea4abdd8854c1de2b0e8d03b7742161))
|
||||
* **auto-fixer:** disable _fix_min_periods for intraday data ([77b0740](https://github.com/TPTBusiness/Predix/commit/77b0740f059349df7e769a378af728aa33b2070e))
|
||||
* **auto-fixer:** fix chained groupby(level=N).groupby('date') pattern ([7d5fe32](https://github.com/TPTBusiness/Predix/commit/7d5fe32b31a19ce8b04bd8f5a430720fdb748f7a))
|
||||
* **auto-fixer:** fix df.loc[instrument] DateParseError on MultiIndex frames ([b7860ea](https://github.com/TPTBusiness/Predix/commit/b7860eafc0ad26384947ce0510ecf4e9f3425807))
|
||||
* **auto-fixer:** fix df['instrument'] KeyError on MultiIndex frames ([aad6bd1](https://github.com/TPTBusiness/Predix/commit/aad6bd1c7c720b3d486e0cf248337f32394773b1))
|
||||
* **auto-fixer:** preserve date dimension in groupby(['instrument','date']) fix ([b58fdd8](https://github.com/TPTBusiness/Predix/commit/b58fdd8be43720b5d4363e0f8de9a01591d4d2dc))
|
||||
* **auto-fixer:** remove ddof from rolling() args, not only from std()/var() ([b0fc328](https://github.com/TPTBusiness/Predix/commit/b0fc328d0d4a041c65d8eeb32cb3f2bb86568406))
|
||||
* **backtest:** replace broken MC permutation test with binomial win-rate test ([c38d894](https://github.com/TPTBusiness/Predix/commit/c38d89478f586825bfca5715a96ca70ccd8791a3))
|
||||
* **factors:** detect and correct look-ahead bias in daily-constant factors ([eb490a4](https://github.com/TPTBusiness/Predix/commit/eb490a461b66cbd815ae53ac5205115754712432))
|
||||
* **factors:** extend look-ahead rules to session factors and add intraday-factor guidance ([c24c100](https://github.com/TPTBusiness/Predix/commit/c24c100442d6487686c0578de0b32d240fcbf215))
|
||||
* **loop:** compress old experiment history in proposal prompt to reduce context size ([4bf90a9](https://github.com/TPTBusiness/Predix/commit/4bf90a905ba8b2aba2a818191c19998088cccaaf))
|
||||
* **loop:** prevent step_idx advance on unhandled exceptions + fix consecutive assistant messages ([5ec4ad1](https://github.com/TPTBusiness/Predix/commit/5ec4ad1b96b5b99ef42bea7bb828cb1ef709a688))
|
||||
|
||||
## [1.3.3](https://github.com/TPTBusiness/Predix/compare/v1.3.2...v1.3.3) (2026-04-25)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **backtest:** replace broken MC permutation test with binomial win-rate test ([c38d894](https://github.com/TPTBusiness/Predix/commit/c38d89478f586825bfca5715a96ca70ccd8791a3))
|
||||
* **factors:** detect and correct look-ahead bias in daily-constant factors ([eb490a4](https://github.com/TPTBusiness/Predix/commit/eb490a461b66cbd815ae53ac5205115754712432))
|
||||
* **factors:** extend look-ahead rules to session factors and add intraday-factor guidance ([c24c100](https://github.com/TPTBusiness/Predix/commit/c24c100442d6487686c0578de0b32d240fcbf215))
|
||||
* **loop:** compress old experiment history in proposal prompt to reduce context size ([4bf90a9](https://github.com/TPTBusiness/Predix/commit/4bf90a905ba8b2aba2a818191c19998088cccaaf))
|
||||
* **strategies:** guard against None IC in acceptance check, disable slow wf_rolling ([2197f52](https://github.com/TPTBusiness/Predix/commit/2197f52150a50ef38d9e70991d7e48c8c30caec4))
|
||||
* **strategies:** handle None ic/sharpe/dd in rejected strategy log output ([ad2ad3a](https://github.com/TPTBusiness/Predix/commit/ad2ad3ab3360ea75ed3bbc90c12098b9c5cc0114))
|
||||
|
||||
## [1.3.2](https://github.com/TPTBusiness/Predix/compare/v1.3.1...v1.3.2) (2026-04-23)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **strategies:** guard against None IC in acceptance check, disable slow wf_rolling ([2197f52](https://github.com/TPTBusiness/Predix/commit/2197f52150a50ef38d9e70991d7e48c8c30caec4))
|
||||
* **strategies:** handle None ic/sharpe/dd in rejected strategy log output ([ad2ad3a](https://github.com/TPTBusiness/Predix/commit/ad2ad3ab3360ea75ed3bbc90c12098b9c5cc0114))
|
||||
|
||||
## [1.3.1](https://github.com/TPTBusiness/Predix/compare/v1.3.0...v1.3.1) (2026-04-21)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **deps:** bump python-dotenv to >=1.2.2 (CVE symlink overwrite) ([126ae7d](https://github.com/TPTBusiness/Predix/commit/126ae7d5fb556b677d09d10221862a0d648d697a))
|
||||
|
||||
## [1.3.0](https://github.com/TPTBusiness/Predix/compare/v1.2.2...v1.3.0) (2026-04-21)
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **backtest:** add rolling walk-forward validation and Monte Carlo trade permutation test ([637a94c](https://github.com/TPTBusiness/Predix/commit/637a94c1d987da763869f4f9b73372a3f37d873c))
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **security:** resolve all 30 Bandit security alerts (B301, B614, B104) ([ce5983d](https://github.com/TPTBusiness/Predix/commit/ce5983d9d59c4c34341fb1ec749e44bbcfc4a1c4))
|
||||
|
||||
## [1.2.2](https://github.com/TPTBusiness/Predix/compare/v1.2.1...v1.2.2) (2026-04-19)
|
||||
|
||||
|
||||
### Documentation
|
||||
|
||||
* **claude:** auto-merge release-please PR after every push ([f500917](https://github.com/TPTBusiness/Predix/commit/f500917b699ee78dc676e84e01574d49bdc8e796))
|
||||
|
||||
## [2.2.0](https://github.com/TPTBusiness/Predix/compare/v2.1.0...v2.2.0) (2026-04-18)
|
||||
|
||||
|
||||
|
||||
@@ -18,6 +18,8 @@ load_dotenv(Path(__file__).parent / ".env")
|
||||
import typer
|
||||
from rich.console import Console
|
||||
|
||||
from rdagent.utils.env import logger
|
||||
|
||||
app = typer.Typer(help="Predix - AI Quantitative Trading Agent")
|
||||
console = Console()
|
||||
|
||||
@@ -510,6 +512,7 @@ def top(
|
||||
if data.get("status") == "success" and data.get("ic") is not None:
|
||||
results.append(data)
|
||||
except Exception:
|
||||
logger.warning("Failed to load factor file %s", f, exc_info=True)
|
||||
continue
|
||||
|
||||
if not results:
|
||||
@@ -659,6 +662,7 @@ def portfolio(
|
||||
if data.get("status") == "success" and data.get("ic") is not None:
|
||||
results.append(data)
|
||||
except Exception:
|
||||
logger.warning("Failed to load factor file %s", f, exc_info=True)
|
||||
continue
|
||||
|
||||
if not results:
|
||||
@@ -956,6 +960,7 @@ def portfolio_simple(
|
||||
if data.get("status") == "success" and data.get("ic") is not None:
|
||||
results.append(data)
|
||||
except Exception:
|
||||
logger.warning("Failed to load factor file %s", f, exc_info=True)
|
||||
continue
|
||||
|
||||
if not results:
|
||||
@@ -1337,6 +1342,7 @@ def build_strategies_ai(
|
||||
if data.get("status") == "success" and data.get("ic") is not None:
|
||||
factors.append(data)
|
||||
except Exception:
|
||||
logger.warning("Failed to load factor file %s", f, exc_info=True)
|
||||
continue
|
||||
|
||||
if len(factors) < 10:
|
||||
@@ -1552,6 +1558,7 @@ def _load_strategies():
|
||||
try:
|
||||
raw = json.loads(p.read_text())
|
||||
except Exception:
|
||||
logger.warning("Failed to load strategy file %s", p, exc_info=True)
|
||||
continue
|
||||
if not isinstance(raw, dict):
|
||||
continue
|
||||
|
||||
+1
-1
@@ -68,7 +68,7 @@ ignore_missing_imports = true
|
||||
module = "llama"
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
addopts = "-l -s --durations=0"
|
||||
addopts = "-l -s --durations=0 -m 'not slow'"
|
||||
log_cli = true
|
||||
log_cli_level = "info"
|
||||
log_date_format = "%Y-%m-%d %H:%M:%S"
|
||||
|
||||
+4
-1
@@ -27,6 +27,8 @@ import typer
|
||||
from rich.console import Console
|
||||
from typing_extensions import Annotated
|
||||
|
||||
from rdagent.utils.env import logger
|
||||
|
||||
from rdagent.app.data_science.loop import main as data_science
|
||||
from rdagent.app.finetune.llm.loop import main as llm_finetune
|
||||
from rdagent.app.general_model.general_model import (
|
||||
@@ -882,6 +884,7 @@ def optimize_portfolio_cli(
|
||||
if data.get("status") == "accepted":
|
||||
strategies.append(data)
|
||||
except Exception:
|
||||
logger.warning("Failed to load strategy file %s", f, exc_info=True)
|
||||
continue
|
||||
|
||||
if not strategies:
|
||||
@@ -1251,7 +1254,7 @@ def start_loop_cli(
|
||||
script_dir = str(Path(__file__).parent.parent.parent.parent)
|
||||
generator = f"python {script_dir}/scripts/predix_smart_strategy_gen.py"
|
||||
logfile = f"{script_dir}/results/logs/generator_loop.log"
|
||||
pidfile = "/tmp/predix_loop.pid"
|
||||
pidfile = "/tmp/predix_loop.pid" # nosec B108 — administrative PID file, single-process daemon
|
||||
|
||||
os.makedirs(f"{script_dir}/results/logs", exist_ok=True)
|
||||
|
||||
|
||||
@@ -201,6 +201,5 @@ class DataScienceBasePropSetting(KaggleBasePropSetting):
|
||||
DS_RD_SETTING = DataScienceBasePropSetting()
|
||||
|
||||
# enable_cross_trace_diversity and llm_select_hypothesis should not be true at the same time
|
||||
assert not (
|
||||
DS_RD_SETTING.enable_cross_trace_diversity and DS_RD_SETTING.llm_select_hypothesis
|
||||
), "enable_cross_trace_diversity and llm_select_hypothesis cannot be true at the same time"
|
||||
if DS_RD_SETTING.enable_cross_trace_diversity and DS_RD_SETTING.llm_select_hypothesis:
|
||||
raise ValueError("enable_cross_trace_diversity and llm_select_hypothesis cannot be true at the same time")
|
||||
|
||||
@@ -58,18 +58,18 @@ def main(
|
||||
|
||||
if user_target_scenario:
|
||||
FT_RD_SETTING.user_target_scenario = user_target_scenario
|
||||
assert (
|
||||
FT_RD_SETTING.user_target_scenario is None
|
||||
), "user_target_scenario is not yet supported, please specify via benchmark and benchmark_description"
|
||||
if FT_RD_SETTING.user_target_scenario is not None:
|
||||
raise ValueError("user_target_scenario is not yet supported, please specify via benchmark and benchmark_description")
|
||||
if upper_data_size_limit:
|
||||
FT_RD_SETTING.upper_data_size_limit = upper_data_size_limit
|
||||
logger.info(f"Set upper_data_size_limit to {FT_RD_SETTING.upper_data_size_limit}")
|
||||
if benchmark and benchmark_description:
|
||||
FT_RD_SETTING.target_benchmark = benchmark
|
||||
FT_RD_SETTING.benchmark_description = benchmark_description
|
||||
assert FT_RD_SETTING.user_target_scenario or (
|
||||
FT_RD_SETTING.target_benchmark and FT_RD_SETTING.benchmark_description
|
||||
), "Either user_target_scenario or target_benchmark must be specified for LLM fine-tuning."
|
||||
if not (
|
||||
FT_RD_SETTING.user_target_scenario or (FT_RD_SETTING.target_benchmark and FT_RD_SETTING.benchmark_description)
|
||||
):
|
||||
raise ValueError("Either user_target_scenario or target_benchmark must be specified for LLM fine-tuning.")
|
||||
|
||||
# Update configuration with provided parameters
|
||||
if dataset:
|
||||
@@ -82,9 +82,8 @@ def main(
|
||||
model_target = FT_RD_SETTING.base_model if FT_RD_SETTING.base_model else "auto selected model"
|
||||
|
||||
# Temporary assertion until auto-selection is implemented
|
||||
assert (
|
||||
FT_RD_SETTING.base_model is not None
|
||||
), "Base model auto selection not yet supported, please specify via --base-model"
|
||||
if FT_RD_SETTING.base_model is None:
|
||||
raise ValueError("Base model auto selection not yet supported, please specify via --base-model")
|
||||
|
||||
logger.info(f"Starting LLM fine-tuning on dataset='{data_set_target}' with model='{model_target}'")
|
||||
|
||||
|
||||
@@ -24,46 +24,12 @@ from rdagent.app.finetune.llm.ui.ft_summary import render_job_summary
|
||||
|
||||
DEFAULT_LOG_BASE = "log/"
|
||||
|
||||
from rdagent.core.utils import safe_resolve_path
|
||||
|
||||
|
||||
def validate_path_within_cwd(user_path: Path) -> Path:
|
||||
"""
|
||||
Validate that a user-provided path is within the current working directory.
|
||||
|
||||
Security: This function prevents path traversal attacks by:
|
||||
1. Resolving the path to its absolute canonical form
|
||||
2. Verifying it's within the CWD boundary using a normalized common prefix
|
||||
3. Rejecting paths outside the boundary with ValueError
|
||||
|
||||
Parameters
|
||||
----------
|
||||
user_path : Path
|
||||
User-provided path to validate
|
||||
|
||||
Returns
|
||||
-------
|
||||
Path
|
||||
Resolved absolute path if valid
|
||||
|
||||
Raises
|
||||
------
|
||||
ValueError
|
||||
If path is outside the current working directory
|
||||
"""
|
||||
safe_root = Path.cwd().resolve()
|
||||
# Expand any user home reference and resolve without requiring the path to exist.
|
||||
resolved_path = user_path.expanduser().resolve(strict=False)
|
||||
|
||||
# Ensure the resolved path is absolute and remains within the safe root.
|
||||
safe_root_str = str(safe_root)
|
||||
resolved_str = str(resolved_path)
|
||||
common = os.path.commonpath([safe_root_str, resolved_str])
|
||||
if common != safe_root_str:
|
||||
raise ValueError("Path is outside the allowed project directory")
|
||||
|
||||
# This will raise ValueError if resolved_path is not within safe_root
|
||||
resolved_path.relative_to(safe_root)
|
||||
|
||||
return resolved_path
|
||||
return safe_resolve_path(user_path, safe_root)
|
||||
|
||||
|
||||
def get_job_options(base_path: Path, safe_root: Path | None = None) -> list[str]:
|
||||
@@ -141,19 +107,14 @@ def main():
|
||||
st.header("Job")
|
||||
base_folder = st.text_input("Base Folder", value=default_log, key="base_folder_input")
|
||||
|
||||
# Normalize and validate the base folder against the configured log root
|
||||
root_real = os.path.realpath(str(Path(default_log).expanduser()))
|
||||
folder_real = os.path.realpath(str(Path(base_folder).expanduser()))
|
||||
if folder_real == root_real or folder_real.startswith(root_real + os.sep):
|
||||
base_path = Path(folder_real)
|
||||
safe_root = Path(root_real)
|
||||
else:
|
||||
safe_root = Path(default_log).expanduser().resolve()
|
||||
try:
|
||||
base_path = safe_resolve_path(Path(base_folder), safe_root)
|
||||
except ValueError:
|
||||
st.error("Invalid base folder: must be within the configured log directory.")
|
||||
safe_root = Path(root_real)
|
||||
base_path = safe_root
|
||||
|
||||
# base_path is validated against safe_root – nosec B614
|
||||
job_options = get_job_options(base_path, safe_root) # nosec B614 – validated above
|
||||
job_options = get_job_options(base_path, safe_root)
|
||||
if job_options:
|
||||
selected_job = st.selectbox("Select Job", job_options, key="job_select")
|
||||
if selected_job.startswith("."):
|
||||
|
||||
@@ -13,6 +13,7 @@ from typing import Any
|
||||
import streamlit as st
|
||||
|
||||
from rdagent.app.finetune.llm.ui.config import EVALUATOR_CONFIG, EventType
|
||||
from rdagent.core.utils import safe_resolve_path
|
||||
from rdagent.log.storage import FileStorage
|
||||
|
||||
|
||||
@@ -89,11 +90,10 @@ def extract_stage(tag: str) -> str:
|
||||
def get_valid_sessions(log_folder: Path, safe_root: Path | None = None) -> list[str]:
|
||||
"""Get list of valid session directories, optionally validating against a safe root."""
|
||||
if safe_root is not None:
|
||||
root_real = os.path.realpath(str(safe_root.expanduser()))
|
||||
folder_real = os.path.realpath(str(log_folder.expanduser()))
|
||||
if not (folder_real == root_real or folder_real.startswith(root_real + os.sep)):
|
||||
try:
|
||||
log_folder = safe_resolve_path(log_folder, safe_root)
|
||||
except ValueError:
|
||||
return []
|
||||
log_folder = Path(folder_real)
|
||||
|
||||
if not log_folder.exists():
|
||||
return []
|
||||
@@ -373,13 +373,11 @@ def parse_event(tag: str, content: Any, timestamp: datetime) -> Event | None:
|
||||
@st.cache_data(ttl=300, hash_funcs={Path: str})
|
||||
def load_ft_session(log_path: Path, safe_root: Path | None = None) -> Session:
|
||||
"""Load events into hierarchical session structure, optionally validating against safe root."""
|
||||
# Validate path is within safe_root if provided
|
||||
if safe_root is not None:
|
||||
root_real = os.path.realpath(str(safe_root.expanduser()))
|
||||
path_real = os.path.realpath(str(log_path.expanduser()))
|
||||
if not (path_real == root_real or path_real.startswith(root_real + os.sep)):
|
||||
try:
|
||||
log_path = safe_resolve_path(log_path, safe_root)
|
||||
except ValueError:
|
||||
return Session()
|
||||
log_path = Path(path_real)
|
||||
|
||||
session = Session()
|
||||
storage = FileStorage(log_path)
|
||||
|
||||
@@ -78,7 +78,8 @@ class QuantRDLoop(RDLoop):
|
||||
while True:
|
||||
if self.get_unfinished_loop_cnt(self.loop_idx) < RD_AGENT_SETTINGS.get_max_parallel():
|
||||
hypo = self._propose()
|
||||
assert hypo.action in ["factor", "model"]
|
||||
if hypo.action not in ["factor", "model"]:
|
||||
raise ValueError(f"hypo.action must be 'factor' or 'model', got {hypo.action!r}")
|
||||
if hypo.action == "factor":
|
||||
exp = self.factor_hypothesis2experiment.convert(hypo, self.trace)
|
||||
else:
|
||||
@@ -322,6 +323,7 @@ class QuantRDLoop(RDLoop):
|
||||
if data.get("status") == "success" and data.get("ic") is not None:
|
||||
factors.append(data)
|
||||
except Exception:
|
||||
logger.warning("Failed to load factor file %s", f, exc_info=True)
|
||||
continue
|
||||
|
||||
if len(factors) < 10:
|
||||
|
||||
@@ -16,55 +16,26 @@ from rdagent.app.rl.ui.components import render_session, render_summary
|
||||
from rdagent.app.rl.ui.config import ALWAYS_VISIBLE_TYPES, OPTIONAL_TYPES
|
||||
from rdagent.app.rl.ui.data_loader import get_summary, get_valid_sessions, load_session
|
||||
from rdagent.app.rl.ui.rl_summary import render_job_summary
|
||||
from rdagent.core.utils import safe_resolve_path
|
||||
|
||||
DEFAULT_LOG_BASE = "log/"
|
||||
|
||||
|
||||
def _safe_resolve(user_input: str | None, safe_root: Path) -> Path:
|
||||
"""
|
||||
Resolve user path relative to safe_root; raise ValueError if it escapes.
|
||||
|
||||
Security: This function prevents path traversal attacks by:
|
||||
1. Rejecting null bytes in user input
|
||||
2. Rejecting Windows drive letters (C:\, D:\, etc.)
|
||||
3. Rejecting absolute paths
|
||||
4. Normalizing path to remove .. traversal attempts
|
||||
5. Validating resolved path is within safe_root using a realpath-based check
|
||||
|
||||
All user-provided paths are validated before filesystem access.
|
||||
"""
|
||||
# Treat the provided safe_root as trusted and canonicalize it once.
|
||||
safe_root = safe_root.expanduser().resolve()
|
||||
|
||||
# Empty input maps to the safe root directory.
|
||||
if not user_input:
|
||||
return safe_root
|
||||
|
||||
# Security check 1: Reject null bytes (path truncation attack)
|
||||
if "\x00" in user_input:
|
||||
raise ValueError("Invalid path: contains null byte")
|
||||
|
||||
try:
|
||||
# Security check 2: Normalize path to resolve .. and . components
|
||||
normalized = os.path.normpath(user_input.strip())
|
||||
|
||||
# Security check 3: Reject Windows drive letters (C:\, D:\, etc.)
|
||||
drive, _ = os.path.splitdrive(normalized)
|
||||
if drive:
|
||||
raise ValueError("Absolute paths with drive letters are not allowed")
|
||||
|
||||
# Security check 4: Reject absolute paths (/, //server/share, etc.)
|
||||
if os.path.isabs(normalized):
|
||||
raise ValueError("Absolute paths are not allowed")
|
||||
|
||||
# Security check 5: Build candidate path under safe_root and fully resolve it.
|
||||
joined = os.path.join(str(safe_root), normalized)
|
||||
resolved_candidate = os.path.realpath(joined)
|
||||
|
||||
# Security check 6: Validate candidate is within safe_root (prevent path traversal)
|
||||
candidate_path = Path(resolved_candidate)
|
||||
# Reconstruct from trusted safe_root so the returned path is root-derived.
|
||||
return safe_root / candidate_path.relative_to(safe_root)
|
||||
joined = safe_root / normalized
|
||||
return safe_resolve_path(joined, safe_root)
|
||||
except (OSError, ValueError) as exc:
|
||||
raise ValueError(f"Invalid path outside of allowed root: {user_input}") from exc
|
||||
|
||||
@@ -82,7 +53,7 @@ def get_job_options(base_path: Path, safe_root: Path | None = None) -> list[str]
|
||||
|
||||
# Security fix: Validate base_path to prevent path traversal
|
||||
try:
|
||||
base_path_resolved = base_path.expanduser().resolve()
|
||||
base_path_resolved = base_path.expanduser().resolve() # nosec B614 — validated against safe_root below via relative_to()
|
||||
|
||||
if safe_root is not None:
|
||||
safe_root_resolved = safe_root.expanduser().resolve()
|
||||
@@ -203,8 +174,7 @@ def main():
|
||||
except ValueError as e:
|
||||
st.warning(str(e))
|
||||
return
|
||||
# job_path is validated by _safe_resolve() above
|
||||
if job_path.exists(): # nosec B614 – path validated by _safe_resolve
|
||||
if job_path.exists():
|
||||
render_job_summary(job_path, safe_root, is_root=is_root_job)
|
||||
else:
|
||||
st.warning(f"Job folder not found: {job_folder}")
|
||||
|
||||
@@ -15,6 +15,7 @@ from typing import Any
|
||||
import streamlit as st
|
||||
|
||||
from rdagent.app.rl.ui.config import EventType
|
||||
from rdagent.core.utils import safe_resolve_path
|
||||
from rdagent.log.storage import FileStorage
|
||||
|
||||
|
||||
@@ -76,11 +77,10 @@ def extract_stage(tag: str) -> str:
|
||||
def get_valid_sessions(log_folder: Path, safe_root: Path | None = None) -> list[str]:
|
||||
"""Get list of valid session directories, optionally validating against a safe root."""
|
||||
if safe_root is not None:
|
||||
root_real = os.path.realpath(str(safe_root.expanduser()))
|
||||
folder_real = os.path.realpath(str(log_folder.expanduser()))
|
||||
if not (folder_real == root_real or folder_real.startswith(root_real + os.sep)):
|
||||
try:
|
||||
log_folder = safe_resolve_path(log_folder, safe_root)
|
||||
except ValueError:
|
||||
return []
|
||||
log_folder = Path(folder_real)
|
||||
|
||||
if not log_folder.exists():
|
||||
return []
|
||||
@@ -245,13 +245,11 @@ def parse_event(tag: str, content: Any, timestamp: datetime) -> Event | None:
|
||||
@st.cache_data(ttl=300, hash_funcs={Path: str})
|
||||
def load_session(log_path: Path, safe_root: Path | None = None) -> Session:
|
||||
"""Load events into hierarchical session structure, optionally validating against safe root."""
|
||||
# Validate path is within safe_root if provided
|
||||
if safe_root is not None:
|
||||
root_real = os.path.realpath(str(safe_root.expanduser()))
|
||||
path_real = os.path.realpath(str(log_path.expanduser()))
|
||||
if not (path_real == root_real or path_real.startswith(root_real + os.sep)):
|
||||
try:
|
||||
log_path = safe_resolve_path(log_path, safe_root)
|
||||
except ValueError:
|
||||
return Session()
|
||||
log_path = Path(path_real)
|
||||
|
||||
session = Session()
|
||||
|
||||
|
||||
@@ -9,6 +9,8 @@ from pathlib import Path
|
||||
import pandas as pd
|
||||
import streamlit as st
|
||||
|
||||
from rdagent.core.utils import safe_resolve_path
|
||||
|
||||
|
||||
def is_valid_task(task_path: Path) -> bool:
|
||||
"""Check if directory is a valid RL task (has __session__ subdirectory)"""
|
||||
@@ -62,14 +64,10 @@ def get_loop_status(task_path: Path, loop_id: int) -> tuple[str, bool | None]:
|
||||
|
||||
|
||||
def _validate_job_path(job_path: Path, safe_root: Path) -> Path:
|
||||
"""Resolve and validate that job_path stays within safe_root."""
|
||||
resolved_root = safe_root.expanduser().resolve()
|
||||
resolved_job = job_path.expanduser().resolve()
|
||||
try:
|
||||
# Reconstruct from trusted root so the returned path is root-derived.
|
||||
return resolved_root / resolved_job.relative_to(resolved_root)
|
||||
return safe_resolve_path(job_path, safe_root)
|
||||
except ValueError:
|
||||
raise ValueError(f"Job path is outside allowed root {resolved_root}")
|
||||
raise ValueError(f"Job path is outside allowed root {safe_root}")
|
||||
|
||||
|
||||
def get_max_loops(job_path: Path, safe_root: Path | None = None) -> int:
|
||||
|
||||
@@ -54,11 +54,11 @@ def rdagent_info():
|
||||
current_version = importlib.metadata.version("rdagent")
|
||||
logger.info(f"RD-Agent version: {current_version}")
|
||||
api_url = f"https://api.github.com/repos/microsoft/RD-Agent/contents/requirements.txt?ref=main"
|
||||
response = requests.get(api_url)
|
||||
response = requests.get(api_url, timeout=30)
|
||||
if response.status_code == 200:
|
||||
files = response.json()
|
||||
file_url = files["download_url"]
|
||||
file_response = requests.get(file_url)
|
||||
file_response = requests.get(file_url, timeout=30)
|
||||
if file_response.status_code == 200:
|
||||
all_file_contents = file_response.text.split("\n")
|
||||
else:
|
||||
|
||||
@@ -5,13 +5,29 @@ from .risk_management import CorrelationAnalyzer, PortfolioOptimizer, AdvancedRi
|
||||
from .vbt_backtest import (
|
||||
DEFAULT_BARS_PER_YEAR,
|
||||
DEFAULT_TXN_COST_BPS,
|
||||
FTMO_INITIAL_CAPITAL,
|
||||
FTMO_MAX_DAILY_LOSS,
|
||||
FTMO_MAX_TOTAL_LOSS,
|
||||
FTMO_MAX_LEVERAGE,
|
||||
FTMO_RISK_PER_TRADE,
|
||||
OOS_START_DEFAULT,
|
||||
WF_IS_YEARS,
|
||||
WF_OOS_YEARS,
|
||||
WF_STEP_YEARS,
|
||||
backtest_from_forward_returns,
|
||||
backtest_signal,
|
||||
backtest_signal_ftmo,
|
||||
monte_carlo_trade_pvalue,
|
||||
walk_forward_rolling,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
'BacktestMetrics', 'FactorBacktester', 'ResultsDatabase',
|
||||
'CorrelationAnalyzer', 'PortfolioOptimizer', 'AdvancedRiskManager',
|
||||
'backtest_signal', 'backtest_from_forward_returns',
|
||||
'backtest_signal', 'backtest_signal_ftmo', 'backtest_from_forward_returns',
|
||||
'monte_carlo_trade_pvalue', 'walk_forward_rolling',
|
||||
'DEFAULT_BARS_PER_YEAR', 'DEFAULT_TXN_COST_BPS',
|
||||
'FTMO_INITIAL_CAPITAL', 'FTMO_MAX_DAILY_LOSS', 'FTMO_MAX_TOTAL_LOSS',
|
||||
'FTMO_MAX_LEVERAGE', 'FTMO_RISK_PER_TRADE', 'OOS_START_DEFAULT',
|
||||
'WF_IS_YEARS', 'WF_OOS_YEARS', 'WF_STEP_YEARS',
|
||||
]
|
||||
|
||||
@@ -71,6 +71,9 @@ class ResultsDatabase:
|
||||
|
||||
self.conn.commit()
|
||||
|
||||
_ALLOWED_TABLES = frozenset({"factors", "backtest_runs", "loop_results"})
|
||||
_ALLOWED_COL_TYPES = frozenset({"REAL", "TEXT", "INTEGER", "BLOB"})
|
||||
|
||||
def _add_column_if_not_exists(self, table: str, column: str, col_type: str) -> None:
|
||||
"""
|
||||
Add a column to a table if it doesn't already exist.
|
||||
@@ -78,20 +81,24 @@ class ResultsDatabase:
|
||||
Parameters
|
||||
----------
|
||||
table : str
|
||||
Table name
|
||||
Table name (must be in _ALLOWED_TABLES)
|
||||
column : str
|
||||
Column name to add
|
||||
Column name to add (alphanumeric + underscore only)
|
||||
col_type : str
|
||||
SQL column type (e.g., 'REAL', 'TEXT')
|
||||
SQL column type (must be in _ALLOWED_COL_TYPES)
|
||||
"""
|
||||
if table not in self._ALLOWED_TABLES:
|
||||
raise ValueError(f"Unknown table: {table!r}")
|
||||
if not column.replace("_", "").isalnum():
|
||||
raise ValueError(f"Invalid column name: {column!r}")
|
||||
if col_type not in self._ALLOWED_COL_TYPES:
|
||||
raise ValueError(f"Invalid column type: {col_type!r}")
|
||||
|
||||
c = self.conn.cursor()
|
||||
try:
|
||||
# Try to query the column - if it fails, it doesn't exist
|
||||
# nosec B608: Internal schema migration, column names are controlled
|
||||
c.execute(f"SELECT {column} FROM {table} LIMIT 1") # nosec B608
|
||||
except sqlite3.OperationalError:
|
||||
# Column doesn't exist, add it
|
||||
c.execute(f"ALTER TABLE {table} ADD COLUMN {column} {col_type}") # nosec B608
|
||||
c.execute("SELECT name FROM pragma_table_info(?)", (table,))
|
||||
existing = {row[0] for row in c.fetchall()}
|
||||
if column not in existing:
|
||||
c.execute(f"ALTER TABLE {table} ADD COLUMN {column} {col_type}")
|
||||
|
||||
def add_factor(self, name: str, type: str = "unknown") -> int:
|
||||
c = self.conn.cursor()
|
||||
@@ -183,16 +190,18 @@ class ResultsDatabase:
|
||||
pd.DataFrame
|
||||
DataFrame with factor names and metrics
|
||||
"""
|
||||
# Map shorthand to full column name
|
||||
_ALLOWED_METRICS = frozenset({
|
||||
'sharpe', 'ic', 'annual_return', 'max_drawdown',
|
||||
'win_rate', 'information_ratio', 'volatility',
|
||||
})
|
||||
metric_map = {
|
||||
'sharpe': 'sharpe',
|
||||
'ic': 'ic',
|
||||
'return': 'annual_return',
|
||||
'drawdown': 'max_drawdown',
|
||||
'win_rate': 'win_rate',
|
||||
'sharpe': 'sharpe', 'ic': 'ic', 'return': 'annual_return',
|
||||
'drawdown': 'max_drawdown', 'win_rate': 'win_rate',
|
||||
'information_ratio': 'information_ratio',
|
||||
}
|
||||
col = metric_map.get(metric, metric)
|
||||
if col not in _ALLOWED_METRICS:
|
||||
raise ValueError(f"Unknown metric: {metric!r}")
|
||||
|
||||
return pd.read_sql_query(
|
||||
f"""SELECT factor_name, ic, sharpe, annual_return, max_drawdown,
|
||||
@@ -201,7 +210,7 @@ class ResultsDatabase:
|
||||
JOIN factors ON factor_id = factors.id
|
||||
WHERE {col} IS NOT NULL
|
||||
ORDER BY {col} DESC
|
||||
LIMIT ?""",
|
||||
LIMIT ?""", # nosec B608 — col is validated against _ALLOWED_METRICS above
|
||||
self.conn,
|
||||
params=[limit]
|
||||
)
|
||||
|
||||
@@ -32,10 +32,22 @@ except ImportError:
|
||||
VBT_AVAILABLE = False
|
||||
|
||||
|
||||
DEFAULT_TXN_COST_BPS = 1.5
|
||||
# 2.35 pip realistic EUR/USD cost: 1.5 spread + 0.5 slippage + 0.35 commission
|
||||
# At EUR/USD ≈ 1.10: 2.35 pip * (0.0001/1.10) ≈ 2.14 bps of notional.
|
||||
DEFAULT_TXN_COST_BPS = 2.14
|
||||
DEFAULT_BARS_PER_YEAR = 252 * 1440 # 252 trading days * 1440 min/day = 362,880
|
||||
EXTREME_BAR_THRESHOLD = 0.05 # |ret| > 5% on a single 1-min bar → suspicious
|
||||
|
||||
# FTMO 100k account rules (enforced in backtest_signal when ftmo=True)
|
||||
FTMO_INITIAL_CAPITAL = 100_000.0
|
||||
FTMO_MAX_DAILY_LOSS = 0.05 # 5% of initial → block new trades rest of day
|
||||
FTMO_MAX_TOTAL_LOSS = 0.10 # 10% of initial → simulation ends
|
||||
# Risk-based position sizing: 0.5% equity risk per trade, 10-pip stop, max 1:30 leverage
|
||||
FTMO_RISK_PER_TRADE = 0.005
|
||||
FTMO_STOP_PIPS = 10
|
||||
FTMO_PIP = 0.0001
|
||||
FTMO_MAX_LEVERAGE = 30
|
||||
|
||||
|
||||
def _compute_trade_pnl(position: pd.Series, strategy_returns: pd.Series) -> pd.Series:
|
||||
"""
|
||||
@@ -259,6 +271,332 @@ def backtest_signal(
|
||||
return result
|
||||
|
||||
|
||||
def _apply_ftmo_mask(
|
||||
signal: pd.Series,
|
||||
close: pd.Series,
|
||||
leverage: float,
|
||||
txn_cost_bps: float,
|
||||
) -> tuple[pd.Series, dict]:
|
||||
"""
|
||||
Apply FTMO daily/total loss rules to a signal series.
|
||||
|
||||
Returns a masked signal (positions zeroed after each limit breach) and
|
||||
a dict of FTMO compliance metrics.
|
||||
"""
|
||||
txn_cost = txn_cost_bps / 10_000.0
|
||||
position = signal.shift(1).fillna(0) * leverage
|
||||
bar_ret = close.pct_change().fillna(0)
|
||||
|
||||
equity = FTMO_INITIAL_CAPITAL
|
||||
peak_day = FTMO_INITIAL_CAPITAL
|
||||
masked = signal.copy()
|
||||
|
||||
daily_breaches = 0
|
||||
total_breached = False
|
||||
total_breach_ts: Optional[pd.Timestamp] = None
|
||||
current_day = None
|
||||
day_start_eq = FTMO_INITIAL_CAPITAL
|
||||
|
||||
pos_prev = 0.0
|
||||
for ts, sig_i in signal.items():
|
||||
day = ts.date() if hasattr(ts, "date") else ts
|
||||
|
||||
if day != current_day:
|
||||
current_day = day
|
||||
day_start_eq = equity
|
||||
|
||||
pos_i = float(signal.at[ts]) * leverage
|
||||
ret_i = float(bar_ret.get(ts, 0.0))
|
||||
cost_i = abs(pos_i - pos_prev) * txn_cost
|
||||
ret_net = pos_prev * ret_i - cost_i
|
||||
equity = equity * (1.0 + ret_net / FTMO_INITIAL_CAPITAL * FTMO_INITIAL_CAPITAL / equity
|
||||
if equity > 0 else 1.0)
|
||||
# Simpler: track as fraction
|
||||
equity += FTMO_INITIAL_CAPITAL * ret_net
|
||||
pos_prev = pos_i
|
||||
|
||||
if total_breached:
|
||||
masked.at[ts] = 0
|
||||
continue
|
||||
|
||||
daily_loss = (equity - day_start_eq) / FTMO_INITIAL_CAPITAL
|
||||
total_loss = (equity - FTMO_INITIAL_CAPITAL) / FTMO_INITIAL_CAPITAL
|
||||
|
||||
if daily_loss < -FTMO_MAX_DAILY_LOSS:
|
||||
daily_breaches += 1
|
||||
day_start_eq = -999 # block rest of day
|
||||
masked.at[ts] = 0
|
||||
|
||||
if total_loss < -FTMO_MAX_TOTAL_LOSS:
|
||||
total_breached = True
|
||||
total_breach_ts = ts
|
||||
masked.at[ts] = 0
|
||||
|
||||
return masked, {
|
||||
"ftmo_daily_breaches": daily_breaches,
|
||||
"ftmo_total_breached": total_breached,
|
||||
"ftmo_total_breach_ts": str(total_breach_ts) if total_breach_ts else None,
|
||||
"ftmo_compliant": not total_breached and daily_breaches == 0,
|
||||
}
|
||||
|
||||
|
||||
OOS_START_DEFAULT = "2024-01-01"
|
||||
|
||||
# Rolling walk-forward default windows (IS years, OOS years, step years)
|
||||
WF_IS_YEARS = 3
|
||||
WF_OOS_YEARS = 1
|
||||
WF_STEP_YEARS = 1
|
||||
|
||||
|
||||
def monte_carlo_trade_pvalue(
|
||||
trade_pnl: pd.Series,
|
||||
n_permutations: int = 1000,
|
||||
seed: int = 0,
|
||||
) -> float:
|
||||
"""
|
||||
Monte Carlo permutation test on trade-level P&L.
|
||||
|
||||
Runs a one-sided binomial test on trade-level win rate.
|
||||
|
||||
Tests H0: win_rate = 0.5 (random trading) against H1: win_rate > 0.5.
|
||||
The ``n_permutations`` parameter is kept for API compatibility but is unused.
|
||||
|
||||
p < 0.05 → win rate is significantly above 50%, indicating a genuine per-trade edge.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
trade_pnl : pd.Series
|
||||
Per-trade net returns (output of ``_compute_trade_pnl``).
|
||||
n_permutations : int
|
||||
Number of random permutations (default 1000).
|
||||
seed : int
|
||||
RNG seed for reproducibility.
|
||||
|
||||
Returns
|
||||
-------
|
||||
float
|
||||
p-value in [0, 1]. Lower is better.
|
||||
"""
|
||||
if len(trade_pnl) < 2:
|
||||
return 1.0
|
||||
trades = trade_pnl.values.copy()
|
||||
# Binomial test: is the win rate significantly above 50%?
|
||||
# p = probability of observing >= n_wins out of n_trades under null (win_rate=0.5).
|
||||
# Low p → strategy has a significant positive edge per trade.
|
||||
from scipy.stats import binomtest
|
||||
n_wins = int((trades > 0).sum())
|
||||
n_total = len(trades)
|
||||
result = binomtest(n_wins, n_total, p=0.5, alternative="greater")
|
||||
return float(result.pvalue)
|
||||
|
||||
|
||||
def walk_forward_rolling(
|
||||
close: pd.Series,
|
||||
signal: pd.Series,
|
||||
leverage: float,
|
||||
txn_cost_bps: float = DEFAULT_TXN_COST_BPS,
|
||||
bars_per_year: int = DEFAULT_BARS_PER_YEAR,
|
||||
is_years: int = WF_IS_YEARS,
|
||||
oos_years: int = WF_OOS_YEARS,
|
||||
step_years: int = WF_STEP_YEARS,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Rolling walk-forward validation: multiple IS/OOS windows shifted by ``step_years``.
|
||||
|
||||
Each window runs an independent FTMO simulation on the IS and OOS slices.
|
||||
Produces aggregate OOS statistics to measure cross-time consistency.
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict with keys:
|
||||
wf_n_windows, wf_oos_sharpe_mean, wf_oos_sharpe_std,
|
||||
wf_oos_monthly_return_mean, wf_oos_consistency (fraction of windows
|
||||
with OOS Sharpe > 0), wf_windows (list of per-window dicts)
|
||||
"""
|
||||
if not isinstance(close.index, pd.DatetimeIndex):
|
||||
return {"wf_n_windows": 0}
|
||||
|
||||
start_year = close.index[0].year
|
||||
end_year = close.index[-1].year
|
||||
|
||||
windows = []
|
||||
yr = start_year
|
||||
while True:
|
||||
is_start = pd.Timestamp(f"{yr}-01-01")
|
||||
is_end = pd.Timestamp(f"{yr + is_years}-01-01")
|
||||
oos_end = pd.Timestamp(f"{yr + is_years + oos_years}-01-01")
|
||||
if oos_end.year > end_year + 1:
|
||||
break
|
||||
is_mask = (close.index >= is_start) & (close.index < is_end)
|
||||
oos_mask = (close.index >= is_end) & (close.index < oos_end)
|
||||
if is_mask.sum() < 1000 or oos_mask.sum() < 1000:
|
||||
yr += step_years
|
||||
continue
|
||||
|
||||
window: Dict[str, Any] = {
|
||||
"is_start": str(is_start.date()),
|
||||
"is_end": str(is_end.date()),
|
||||
"oos_start": str(is_end.date()),
|
||||
"oos_end": str(oos_end.date()),
|
||||
}
|
||||
for mask, prefix in [(is_mask, "is"), (oos_mask, "oos")]:
|
||||
close_s = close.loc[mask]
|
||||
signal_s = signal.loc[mask]
|
||||
masked_s, _ = _apply_ftmo_mask(signal_s, close_s, leverage, txn_cost_bps)
|
||||
r = backtest_signal(close=close_s, signal=masked_s,
|
||||
txn_cost_bps=txn_cost_bps, bars_per_year=bars_per_year)
|
||||
window[f"{prefix}_sharpe"] = r.get("sharpe", 0.0)
|
||||
window[f"{prefix}_monthly_return_pct"] = r.get("monthly_return_pct", 0.0)
|
||||
window[f"{prefix}_n_trades"] = r.get("n_trades", 0)
|
||||
windows.append(window)
|
||||
yr += step_years
|
||||
|
||||
if not windows:
|
||||
return {"wf_n_windows": 0}
|
||||
|
||||
oos_sharpes = [w["oos_sharpe"] for w in windows]
|
||||
oos_monthly = [w["oos_monthly_return_pct"] for w in windows]
|
||||
return {
|
||||
"wf_n_windows": len(windows),
|
||||
"wf_oos_sharpe_mean": float(np.mean(oos_sharpes)),
|
||||
"wf_oos_sharpe_std": float(np.std(oos_sharpes)),
|
||||
"wf_oos_monthly_return_mean": float(np.mean(oos_monthly)),
|
||||
"wf_oos_consistency": float(np.mean([s > 0 for s in oos_sharpes])),
|
||||
"wf_windows": windows,
|
||||
}
|
||||
|
||||
|
||||
def backtest_signal_ftmo(
|
||||
close: pd.Series,
|
||||
signal: pd.Series,
|
||||
txn_cost_bps: float = DEFAULT_TXN_COST_BPS,
|
||||
eurusd_price: float = 1.10,
|
||||
risk_pct: float = FTMO_RISK_PER_TRADE,
|
||||
stop_pips: float = FTMO_STOP_PIPS,
|
||||
max_leverage: float = FTMO_MAX_LEVERAGE,
|
||||
bars_per_year: int = DEFAULT_BARS_PER_YEAR,
|
||||
forward_returns: Optional[pd.Series] = None,
|
||||
oos_start: Optional[str] = OOS_START_DEFAULT,
|
||||
wf_rolling: bool = False,
|
||||
mc_n_permutations: int = 0,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
FTMO-compliant backtest of a strategy signal on EUR/USD.
|
||||
|
||||
Applies on top of ``backtest_signal``:
|
||||
- Realistic costs: default 2.14 bps (≈ 2.35 pip spread+slippage+commission)
|
||||
- Risk-based position sizing: risk_pct equity per trade, stop_pips hard stop
|
||||
- Max leverage cap: max_leverage (default 1:30, FTMO standard)
|
||||
- FTMO daily loss limit (5%): positions zeroed rest of day after breach
|
||||
- FTMO total loss limit (10%): all positions zeroed after breach
|
||||
- FTMO-specific metrics added to result dict
|
||||
- Walk-forward OOS split: IS metrics (before oos_start) + OOS metrics (after)
|
||||
|
||||
Parameters
|
||||
----------
|
||||
close : pd.Series
|
||||
1-min EUR/USD close prices.
|
||||
signal : pd.Series
|
||||
Raw strategy signal in {-1, 0, +1}.
|
||||
txn_cost_bps : float
|
||||
Transaction cost in bps (default 2.14 ≈ 2.35 pip on EUR/USD).
|
||||
eurusd_price : float
|
||||
Representative EUR/USD price for pip→bps conversion (default 1.10).
|
||||
risk_pct : float
|
||||
Fraction of equity risked per trade (default 0.005 = 0.5%).
|
||||
stop_pips : float
|
||||
Hard stop-loss distance in pips (default 10).
|
||||
max_leverage : float
|
||||
Maximum leverage (default 30 = FTMO 1:30).
|
||||
oos_start : str or None
|
||||
Start of out-of-sample period (ISO date). None disables OOS split.
|
||||
wf_rolling : bool
|
||||
If True, run rolling walk-forward validation (multiple IS/OOS windows).
|
||||
Results are stored under ``wf_*`` keys. Default False.
|
||||
mc_n_permutations : int
|
||||
Number of Monte Carlo trade permutations. 0 = disabled (default).
|
||||
When > 0, computes ``mc_pvalue``: fraction of permuted sequences whose
|
||||
total return >= real total return. p < 0.05 indicates a genuine edge.
|
||||
"""
|
||||
stop_price = stop_pips * FTMO_PIP
|
||||
leverage_by_risk = risk_pct / (stop_price / eurusd_price)
|
||||
leverage = min(leverage_by_risk, max_leverage)
|
||||
|
||||
masked_signal, ftmo_metrics = _apply_ftmo_mask(signal, close, leverage, txn_cost_bps)
|
||||
|
||||
result = backtest_signal(
|
||||
close=close,
|
||||
signal=masked_signal,
|
||||
txn_cost_bps=txn_cost_bps,
|
||||
bars_per_year=bars_per_year,
|
||||
forward_returns=forward_returns,
|
||||
)
|
||||
|
||||
result.update(ftmo_metrics)
|
||||
result["ftmo_leverage"] = round(leverage, 2)
|
||||
result["ftmo_risk_pct"] = risk_pct
|
||||
result["ftmo_stop_pips"] = stop_pips
|
||||
|
||||
# Re-scale reported equity metrics to FTMO_INITIAL_CAPITAL
|
||||
result["ftmo_end_equity"] = FTMO_INITIAL_CAPITAL * (1 + result.get("total_return", 0))
|
||||
result["ftmo_monthly_profit"] = FTMO_INITIAL_CAPITAL * result.get("monthly_return", 0)
|
||||
|
||||
# Walk-forward OOS split
|
||||
if oos_start is not None:
|
||||
oos_ts = pd.Timestamp(oos_start)
|
||||
is_mask = close.index < oos_ts
|
||||
oos_mask = close.index >= oos_ts
|
||||
|
||||
def _split_bt(mask: "pd.Series[bool]", prefix: str) -> None:
|
||||
if mask.sum() < 100:
|
||||
return
|
||||
close_s = close.loc[mask]
|
||||
signal_s = signal.loc[mask] # raw signal, not masked — fresh FTMO sim per period
|
||||
fwd_split = forward_returns.loc[mask] if forward_returns is not None else None
|
||||
masked_s, _ = _apply_ftmo_mask(signal_s, close_s, leverage, txn_cost_bps)
|
||||
split_result = backtest_signal(
|
||||
close=close_s,
|
||||
signal=masked_s,
|
||||
txn_cost_bps=txn_cost_bps,
|
||||
bars_per_year=bars_per_year,
|
||||
forward_returns=fwd_split,
|
||||
)
|
||||
for k, v in split_result.items():
|
||||
if k not in ("equity_curve", "status"):
|
||||
result[f"{prefix}_{k}"] = v
|
||||
|
||||
_split_bt(is_mask, "is")
|
||||
_split_bt(oos_mask, "oos")
|
||||
|
||||
result["oos_start"] = oos_start
|
||||
result["is_n_bars"] = int(is_mask.sum())
|
||||
result["oos_n_bars"] = int(oos_mask.sum())
|
||||
|
||||
# Rolling walk-forward validation
|
||||
if wf_rolling:
|
||||
wf = walk_forward_rolling(
|
||||
close=close,
|
||||
signal=signal,
|
||||
leverage=leverage,
|
||||
txn_cost_bps=txn_cost_bps,
|
||||
bars_per_year=bars_per_year,
|
||||
)
|
||||
result.update(wf)
|
||||
|
||||
# Monte Carlo trade permutation test
|
||||
if mc_n_permutations > 0:
|
||||
position = masked_signal.shift(1).fillna(0)
|
||||
bar_ret = close.pct_change().fillna(0)
|
||||
txn_cost = txn_cost_bps / 10_000.0
|
||||
position_change = position.diff().abs().fillna(position.abs())
|
||||
strat_ret = position * bar_ret - position_change * txn_cost
|
||||
trade_pnl = _compute_trade_pnl(position, strat_ret)
|
||||
result["mc_pvalue"] = monte_carlo_trade_pvalue(trade_pnl, mc_n_permutations)
|
||||
result["mc_n_permutations"] = mc_n_permutations
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def backtest_from_forward_returns(
|
||||
factor_values: pd.Series,
|
||||
forward_returns: pd.Series,
|
||||
|
||||
@@ -75,8 +75,10 @@ class CoSTEER(Developer[Experiment]):
|
||||
|
||||
def _get_last_fb(self) -> CoSTEERMultiFeedback:
|
||||
fb = self.evolve_agent.evolving_trace[-1].feedback
|
||||
assert fb is not None, "feedback is None"
|
||||
assert isinstance(fb, CoSTEERMultiFeedback), "feedback must be of type CoSTEERMultiFeedback"
|
||||
if fb is None:
|
||||
raise AssertionError("feedback is None")
|
||||
if not isinstance(fb, CoSTEERMultiFeedback):
|
||||
raise TypeError("feedback must be of type CoSTEERMultiFeedback")
|
||||
return fb
|
||||
|
||||
def should_use_new_evo(self, base_fb: CoSTEERMultiFeedback | None, new_fb: CoSTEERMultiFeedback) -> bool:
|
||||
@@ -121,7 +123,8 @@ class CoSTEER(Developer[Experiment]):
|
||||
|
||||
for evo_exp in self.evolve_agent.multistep_evolve(evo_exp, self.evaluator):
|
||||
iteration_count += 1
|
||||
assert isinstance(evo_exp, Experiment) # multiple inheritance
|
||||
if not isinstance(evo_exp, Experiment):
|
||||
raise TypeError("evo_exp must be an instance of Experiment")
|
||||
evo_fb = self._get_last_fb()
|
||||
update_fallback = self.should_use_new_evo(
|
||||
base_fb=fallback_evo_fb,
|
||||
@@ -154,7 +157,8 @@ class CoSTEER(Developer[Experiment]):
|
||||
evo_exp = fallback_evo_exp
|
||||
evo_exp.recover_ws_ckp()
|
||||
evo_fb = fallback_evo_fb
|
||||
assert evo_fb is not None # multistep_evolve should run at least once
|
||||
if evo_fb is None:
|
||||
raise AssertionError("multistep_evolve should run at least once")
|
||||
evo_exp = self._exp_postprocess_by_feedback(evo_exp, evo_fb)
|
||||
except CoderError as e:
|
||||
e.caused_by_timeout = reached_max_seconds
|
||||
@@ -264,9 +268,12 @@ class CoSTEER(Developer[Experiment]):
|
||||
- Raise Error if it failed to handle the develop task
|
||||
-
|
||||
"""
|
||||
assert isinstance(evo, Experiment)
|
||||
assert isinstance(feedback, CoSTEERMultiFeedback)
|
||||
assert len(evo.sub_workspace_list) == len(feedback)
|
||||
if not isinstance(evo, Experiment):
|
||||
raise TypeError("evo must be an instance of Experiment")
|
||||
if not isinstance(feedback, CoSTEERMultiFeedback):
|
||||
raise TypeError("feedback must be an instance of CoSTEERMultiFeedback")
|
||||
if len(evo.sub_workspace_list) != len(feedback):
|
||||
raise ValueError("Length of sub_workspace_list must match length of feedback")
|
||||
|
||||
# FIXME: when whould the feedback be None?
|
||||
failed_feedbacks = [
|
||||
|
||||
@@ -122,7 +122,8 @@ class MultiProcessEvolvingStrategy(EvolvingStrategy):
|
||||
last_feedback = None
|
||||
if len(evolving_trace) > 0:
|
||||
last_feedback = evolving_trace[-1].feedback
|
||||
assert isinstance(last_feedback, CoSTEERMultiFeedback)
|
||||
if not isinstance(last_feedback, CoSTEERMultiFeedback):
|
||||
raise TypeError("last_feedback must be of type CoSTEERMultiFeedback")
|
||||
|
||||
# 1.找出需要evolve的task
|
||||
to_be_finished_task_index: list[int] = []
|
||||
|
||||
@@ -1028,7 +1028,8 @@ class CoSTEERKnowledgeBaseV2(EvolvingKnowledgeBase):
|
||||
|
||||
"""
|
||||
node_count = len(nodes)
|
||||
assert node_count >= 2, "nodes length must >=2"
|
||||
if node_count < 2:
|
||||
raise ValueError("nodes length must >=2")
|
||||
intersection_node_list = []
|
||||
if output_intersection_origin:
|
||||
origin_list = []
|
||||
|
||||
@@ -54,7 +54,8 @@ def get_ds_env(
|
||||
ValueError: If the env_type is not recognized.
|
||||
"""
|
||||
conf = DSCoderCoSTEERSettings()
|
||||
assert conf_type in ["kaggle", "mlebench"], f"Unknown conf_type: {conf_type}"
|
||||
if conf_type not in ["kaggle", "mlebench"]:
|
||||
raise ValueError(f"Unknown conf_type: {conf_type}")
|
||||
|
||||
if conf.env_type == "docker":
|
||||
env_conf = DSDockerConf() if conf_type == "kaggle" else MLEBDockerConf()
|
||||
@@ -79,7 +80,8 @@ def get_clear_ws_cmd(stage: Literal["before_training", "before_inference"] = "be
|
||||
"""
|
||||
Clean the files in workspace to a specific stage
|
||||
"""
|
||||
assert stage in ["before_training", "before_inference"], f"Unknown stage: {stage}"
|
||||
if stage not in ["before_training", "before_inference"]:
|
||||
raise ValueError(f"Unknown stage: {stage}")
|
||||
if DS_RD_SETTING.enable_model_dump and stage == "before_training":
|
||||
cmd = "rm -r submission.csv scores.csv models trace.log"
|
||||
else:
|
||||
|
||||
@@ -13,7 +13,7 @@ File structure
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from jinja2 import Environment, StrictUndefined
|
||||
from jinja2 import Environment, StrictUndefined, select_autoescape
|
||||
|
||||
from rdagent.app.data_science.conf import DS_RD_SETTING
|
||||
from rdagent.components.coder.CoSTEER.evaluators import (
|
||||
@@ -88,7 +88,7 @@ class EnsembleMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
|
||||
code_spec = workspace.file_dict["spec/ensemble.md"]
|
||||
else:
|
||||
test_code = (
|
||||
Environment(undefined=StrictUndefined)
|
||||
Environment(undefined=StrictUndefined, autoescape=select_autoescape())
|
||||
.from_string((DIRNAME / "eval_tests" / "ensemble_test.txt").read_text())
|
||||
.render(
|
||||
model_names=[
|
||||
|
||||
@@ -2,7 +2,7 @@ import json
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
from jinja2 import Environment, StrictUndefined
|
||||
from jinja2 import Environment, StrictUndefined, select_autoescape
|
||||
|
||||
from rdagent.app.data_science.conf import DS_RD_SETTING
|
||||
from rdagent.components.coder.CoSTEER.evaluators import (
|
||||
@@ -55,7 +55,7 @@ class EnsembleCoSTEEREvaluator(CoSTEEREvaluator):
|
||||
fname = "test/ensemble_test.txt"
|
||||
test_code = (DIRNAME / "eval_tests" / "ensemble_test.txt").read_text()
|
||||
test_code = (
|
||||
Environment(undefined=StrictUndefined)
|
||||
Environment(undefined=StrictUndefined, autoescape=select_autoescape())
|
||||
.from_string(test_code)
|
||||
.render(
|
||||
model_names=[
|
||||
|
||||
@@ -51,13 +51,23 @@ class FactorAutoFixer:
|
||||
self.fixes_applied = []
|
||||
fixed_code = code
|
||||
|
||||
# Apply fixes in order - groupby fixes MUST come before min_periods fixes
|
||||
# Apply fixes in order
|
||||
# NOTE: _fix_min_periods is intentionally excluded — it increased min_periods to
|
||||
# match window size, which causes all-NaN output for intraday data with 96 bars/day
|
||||
# (window=240 > 96 means zero valid bars per day). The LLM sets its own min_periods.
|
||||
fix_methods = [
|
||||
self._fix_groupby_apply_to_transform, # First: fix groupby patterns
|
||||
self._fix_min_periods, # Second: fix min_periods in resulting rolling calls
|
||||
self._fix_inf_nan_handling, # Third: add inf/nan handling
|
||||
self._fix_data_range_processing, # Fourth: ensure full data range
|
||||
self._fix_multiindex_groupby, # Fifth: ensure groupby on MultiIndex
|
||||
self._fix_instrument_column_access, # First: fix df['instrument'] on MultiIndex
|
||||
self._fix_instrument_loc_multiindex, # Second: fix df.loc[instrument_var] on MultiIndex
|
||||
self._fix_zero_volume_proxy, # Third: replace zero $volume with range proxy
|
||||
self._fix_reset_index_groupby, # Fourth: fix groupby(level=N) after reset_index()
|
||||
self._fix_groupby_mixed_levels, # Fifth: fix groupby(level=[int, str])
|
||||
self._fix_groupby_column_on_multiindex, # Sixth: fix groupby(['instrument','date']) on MultiIndex
|
||||
self._fix_chained_groupby, # Seventh: fix groupby(level=N).groupby('date') chain
|
||||
self._fix_rolling_ddof, # Eighth: remove unsupported ddof kwarg
|
||||
self._fix_groupby_apply_to_transform, # Ninth: fix groupby patterns
|
||||
self._fix_inf_nan_handling, # Tenth: add inf/nan handling
|
||||
self._fix_data_range_processing, # Eleventh: ensure full data range
|
||||
self._fix_multiindex_groupby, # Twelfth: ensure groupby on MultiIndex
|
||||
]
|
||||
|
||||
for fix_method in fix_methods:
|
||||
@@ -75,6 +85,352 @@ class FactorAutoFixer:
|
||||
|
||||
return fixed_code
|
||||
|
||||
def _fix_instrument_column_access(self, code: str) -> str:
|
||||
"""
|
||||
Fix: df['instrument'] raises KeyError on a MultiIndex DataFrame because
|
||||
'instrument' is an index level (level 1), not a column.
|
||||
|
||||
Replace df['instrument'] with df.index.get_level_values('instrument')
|
||||
but only when the DataFrame has a MultiIndex (not after reset_index which
|
||||
would have promoted it to a real column).
|
||||
|
||||
Also fixes df.reset_index()['instrument'] correctly since after reset_index
|
||||
the column exists.
|
||||
"""
|
||||
fixed_code = code
|
||||
|
||||
# Skip if already fixed or if reset_index() is being used before the access
|
||||
# We only fix bare df['instrument'] where df is the original MultiIndex frame.
|
||||
# Heuristic: if the assignment lhs or context shows reset_index, leave it alone.
|
||||
|
||||
# Pattern: <varname>['instrument'] where varname is NOT a reset_index result
|
||||
reset_vars = set(re.findall(r'(\w+)\s*=\s*\w[^=\n]*\.reset_index\(', fixed_code))
|
||||
|
||||
def _replace_instrument_access(m: re.Match) -> str:
|
||||
var = m.group(1)
|
||||
if var in reset_vars:
|
||||
return m.group(0) # leave reset_index vars alone — column exists
|
||||
self.fixes_applied.append(f"instrument_column: {var}['instrument'] → get_level_values(1)")
|
||||
return f"{var}.index.get_level_values(1)"
|
||||
|
||||
# Exclude assignment targets: var['instrument'] = ... must not become
|
||||
# var.index.get_level_values(1) = ... (SyntaxError: cannot assign to function call)
|
||||
fixed_code = re.sub(r"(\w+)\['instrument'\](?!\s*=)", _replace_instrument_access, fixed_code)
|
||||
|
||||
return fixed_code
|
||||
|
||||
def _fix_instrument_loc_multiindex(self, code: str) -> str:
|
||||
"""
|
||||
Fix: df.loc[instrument_var] raises DateParseError on a (datetime, instrument)
|
||||
MultiIndex because pandas tries to match the instrument string against the
|
||||
datetime level (level 0).
|
||||
|
||||
Pattern detected: for-loops iterating over get_level_values('instrument') or
|
||||
get_level_values(1) where the loop variable is then used as df.loc[loop_var].
|
||||
|
||||
Replacement: df.loc[instrument_var] → df.xs(instrument_var, level=1)
|
||||
"""
|
||||
fixed_code = code
|
||||
|
||||
# Find variables iterated from get_level_values('instrument') or get_level_values(1)
|
||||
inst_vars = set(
|
||||
re.findall(
|
||||
r"for\s+(\w+)\s+in\s+.+?\.get_level_values\s*\(\s*(?:1|['\"]instrument['\"])\s*\)[^:\n]*:",
|
||||
code,
|
||||
)
|
||||
)
|
||||
|
||||
if not inst_vars:
|
||||
return fixed_code
|
||||
|
||||
for var in inst_vars:
|
||||
# Replace DF.loc[var] (read) with DF.xs(var, level=1)
|
||||
# Exclude write-back patterns (DF.loc[var] = ...) — leave those as-is
|
||||
def _make_replacer(v: str):
|
||||
def _replace(m: re.Match) -> str:
|
||||
df_var = m.group(1)
|
||||
self.fixes_applied.append(
|
||||
f"instrument_loc: {df_var}.loc[{v}] → {df_var}.xs({v}, level=1)"
|
||||
)
|
||||
return f"{df_var}.xs({v}, level=1)"
|
||||
|
||||
return _replace
|
||||
|
||||
# Only match when NOT followed by ' =' (assignment)
|
||||
fixed_code = re.sub(
|
||||
rf"(\w+)\.loc\[\s*{re.escape(var)}\s*\](?!\s*=)",
|
||||
_make_replacer(var),
|
||||
fixed_code,
|
||||
)
|
||||
|
||||
return fixed_code
|
||||
|
||||
def _fix_zero_volume_proxy(self, code: str) -> str:
|
||||
"""
|
||||
Fix: $volume is always 0 in our EUR/USD dataset (FX has no real volume).
|
||||
Any factor using $volume (VWAP, volume-weighted returns, etc.) produces
|
||||
all-NaN output because 0*price=0 and sum(0)/sum(0)=NaN.
|
||||
|
||||
Insert a guard right after pd.read_hdf() that replaces zero volume with
|
||||
the intraday price-range proxy ($high - $low) so volume-weighted factors
|
||||
produce meaningful signals.
|
||||
"""
|
||||
if "'$volume'" not in code and '"$volume"' not in code:
|
||||
return code
|
||||
|
||||
# Already patched
|
||||
if "volume proxy" in code:
|
||||
return code
|
||||
|
||||
lines = code.splitlines()
|
||||
insert_after = -1
|
||||
df_var = "df"
|
||||
indent = " "
|
||||
|
||||
for i, line in enumerate(lines):
|
||||
if "read_hdf(" in line:
|
||||
m = re.match(r"(\s*)(\w+)\s*=\s*", line)
|
||||
if m:
|
||||
indent = m.group(1)
|
||||
df_var = m.group(2)
|
||||
else:
|
||||
m2 = re.match(r"(\s*)", line)
|
||||
indent = m2.group(1) if m2 else " "
|
||||
insert_after = i
|
||||
break
|
||||
|
||||
if insert_after == -1:
|
||||
return code
|
||||
|
||||
proxy_lines = [
|
||||
f"{indent}# volume proxy: $volume is always 0 in FX data — use price-range as proxy",
|
||||
f"{indent}if ({df_var}['$volume'] == 0).all():",
|
||||
f"{indent} {df_var}['$volume'] = {df_var}['$high'] - {df_var}['$low']",
|
||||
]
|
||||
lines = lines[: insert_after + 1] + proxy_lines + lines[insert_after + 1 :]
|
||||
self.fixes_applied.append("volume_proxy: replaced zero $volume with ($high - $low)")
|
||||
return "\n".join(lines)
|
||||
|
||||
def _fix_reset_index_groupby(self, code: str) -> str:
|
||||
"""
|
||||
Fix: groupby(level=N) on a variable created by .reset_index() fails because
|
||||
reset_index() converts the MultiIndex into regular columns, leaving a plain
|
||||
RangeIndex. Replace groupby(level=N) on such variables with
|
||||
groupby('instrument').
|
||||
|
||||
Detected pattern:
|
||||
varname = <anything>.reset_index(...)
|
||||
...
|
||||
varname.groupby(level=0|1)
|
||||
"""
|
||||
fixed_code = code
|
||||
|
||||
# Find all variables assigned via reset_index()
|
||||
reset_vars = set(re.findall(r'(\w+)\s*=\s*\w[^=\n]*\.reset_index\(', fixed_code))
|
||||
|
||||
for var in reset_vars:
|
||||
# Replace var.groupby(level=N) with var.groupby('instrument')
|
||||
pattern = rf'{re.escape(var)}\.groupby\(level\s*=\s*\d+\)'
|
||||
if re.search(pattern, fixed_code):
|
||||
fixed_code = re.sub(pattern, f"{var}.groupby('instrument')", fixed_code)
|
||||
self.fixes_applied.append(f"reset_index_groupby: {var}.groupby(level=N) → groupby('instrument')")
|
||||
|
||||
return fixed_code
|
||||
|
||||
def _fix_groupby_mixed_levels(self, code: str) -> str:
|
||||
"""
|
||||
Fix: groupby(level=[int, 'str']) raises AssertionError because string level
|
||||
names don't exist on an unnamed MultiIndex. Keep only integer levels.
|
||||
|
||||
Pattern: .groupby(level=[0, 'date']) → .groupby(level=0)
|
||||
.groupby(level=[1, 'date']) → .groupby(level=1)
|
||||
"""
|
||||
fixed_code = code
|
||||
|
||||
def _keep_int_levels(m):
|
||||
inner = m.group(1)
|
||||
ints = re.findall(r'\b(\d+)\b', inner)
|
||||
if not ints:
|
||||
return m.group(0)
|
||||
replacement = f'.groupby(level={ints[0]})' if len(ints) == 1 else f'.groupby(level=[{", ".join(ints)}])'
|
||||
self.fixes_applied.append(f"mixed_levels: groupby(level=[...,str]) → {replacement}")
|
||||
return replacement
|
||||
|
||||
fixed_code = re.sub(r'\.groupby\(level=\[([^\]]+)\]\)', _keep_int_levels, fixed_code)
|
||||
return fixed_code
|
||||
|
||||
def _fix_groupby_column_on_multiindex(self, code: str) -> str:
|
||||
"""
|
||||
Fix: groupby(['instrument', 'date']) on a MultiIndex (datetime, instrument)
|
||||
DataFrame fails with KeyError because those are index levels, not columns.
|
||||
|
||||
Correct replacement preserves BOTH dimensions so intraday calculations reset
|
||||
per day:
|
||||
var.groupby(['instrument', 'date'])
|
||||
→ var.groupby([var.index.get_level_values(1), var.index.get_level_values(0).normalize()])
|
||||
|
||||
Single-column groupby(['instrument']) is correctly replaced with groupby(level=1).
|
||||
Note: do NOT convert groupby('instrument') → groupby(level=1) here — that would
|
||||
undo the reset_index_groupby fix which correctly emits groupby('instrument').
|
||||
"""
|
||||
fixed_code = code
|
||||
|
||||
# Variables created via reset_index() have a plain RangeIndex — applying
|
||||
# get_level_values() on them would raise AttributeError. Skip those.
|
||||
reset_vars = set(re.findall(r'(\w+)\s*=\s*\w[^=\n]*\.reset_index\(', fixed_code))
|
||||
|
||||
def _replace_two_col_groupby(m: re.Match, order: str) -> str:
|
||||
var = m.group(1)
|
||||
if var in reset_vars:
|
||||
return m.group(0) # leave reset_index vars alone — RangeIndex, not MultiIndex
|
||||
if order == "instrument_date":
|
||||
repl = (
|
||||
f"{var}.groupby([{var}.index.get_level_values(1), "
|
||||
f"{var}.index.get_level_values(0).normalize()])"
|
||||
)
|
||||
else: # date_instrument
|
||||
repl = (
|
||||
f"{var}.groupby([{var}.index.get_level_values(0).normalize(), "
|
||||
f"{var}.index.get_level_values(1)])"
|
||||
)
|
||||
self.fixes_applied.append(f"multiindex_groupby: {m.group(0)[:60]} → two-level")
|
||||
return repl
|
||||
|
||||
# groupby(['instrument', 'date']) — capture variable name before .groupby
|
||||
fixed_code = re.sub(
|
||||
r'(\w+)\.groupby\(\[\'instrument\',\s*\'date\'\]\)',
|
||||
lambda m: _replace_two_col_groupby(m, "instrument_date"),
|
||||
fixed_code,
|
||||
)
|
||||
# groupby(['date', 'instrument'])
|
||||
fixed_code = re.sub(
|
||||
r'(\w+)\.groupby\(\[\'date\',\s*\'instrument\'\]\)',
|
||||
lambda m: _replace_two_col_groupby(m, "date_instrument"),
|
||||
fixed_code,
|
||||
)
|
||||
# single: groupby(['instrument']) → groupby(level=1), but not on reset_index vars
|
||||
def _replace_single_instrument_groupby(m: re.Match) -> str:
|
||||
# Look backwards to find the variable name
|
||||
prefix = fixed_code[: m.start()]
|
||||
var_match = re.search(r'(\w+)\s*$', prefix)
|
||||
var = var_match.group(1) if var_match else ''
|
||||
if var in reset_vars:
|
||||
return m.group(0)
|
||||
self.fixes_applied.append("multiindex_groupby: groupby(['instrument']) → groupby(level=1)")
|
||||
return ".groupby(level=1)"
|
||||
|
||||
if re.search(r"\.groupby\(\['instrument'\]\)", fixed_code):
|
||||
fixed_code = re.sub(r"\.groupby\(\['instrument'\]\)", _replace_single_instrument_groupby, fixed_code)
|
||||
|
||||
# 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:
|
||||
"""
|
||||
Fix two broken patterns the LLM generates when trying to group by (instrument, date):
|
||||
|
||||
Pattern A — chained groupby (runtime AttributeError):
|
||||
var.groupby(level=1).groupby('date')
|
||||
→ var.groupby([var.index.get_level_values(1),
|
||||
var.index.get_level_values(0).normalize()])
|
||||
|
||||
Pattern B — keyword arg inside list (SyntaxError):
|
||||
var.groupby([level=1, 'date'])
|
||||
→ same two-level replacement
|
||||
"""
|
||||
fixed_code = code
|
||||
|
||||
def _two_level(var: str, tag: str) -> str:
|
||||
self.fixes_applied.append(f"chained_groupby: {tag} → two-level")
|
||||
return (
|
||||
f"{var}.groupby([{var}.index.get_level_values(1), "
|
||||
f"{var}.index.get_level_values(0).normalize()])"
|
||||
)
|
||||
|
||||
# Pattern A: var.groupby(level=N).groupby('date')
|
||||
fixed_code = re.sub(
|
||||
r'(\w+)\.groupby\(level=\d+\)\.groupby\(["\']date["\']\)',
|
||||
lambda m: _two_level(m.group(1), m.group(0)[:60]),
|
||||
fixed_code,
|
||||
)
|
||||
|
||||
# Pattern B: .groupby([level=N, 'date']) — SyntaxError in Python.
|
||||
# The variable before .groupby may be complex (e.g. df[mask]) so we don't
|
||||
# try to capture it; we use df as the index reference (always correct since
|
||||
# all filtered frames share df's MultiIndex structure).
|
||||
def _two_level_df(tag: str) -> str:
|
||||
self.fixes_applied.append(f"chained_groupby: {tag} → two-level")
|
||||
return ".groupby([df.index.get_level_values(1), df.index.get_level_values(0).normalize()])"
|
||||
|
||||
fixed_code = re.sub(
|
||||
r'\.groupby\(\[\s*level\s*=\s*\d+\s*,\s*["\']?date["\']?\s*\]\)',
|
||||
lambda m: _two_level_df(m.group(0)[:60]),
|
||||
fixed_code,
|
||||
)
|
||||
# Also handle reversed order: ['date', level=N]
|
||||
fixed_code = re.sub(
|
||||
r'\.groupby\(\[\s*["\']?date["\']?\s*,\s*level\s*=\s*\d+\s*\]\)',
|
||||
lambda m: _two_level_df(m.group(0)[:60]),
|
||||
fixed_code,
|
||||
)
|
||||
|
||||
return fixed_code
|
||||
|
||||
def _fix_rolling_ddof(self, code: str) -> str:
|
||||
"""
|
||||
Fix: pandas rolling() does not accept a ddof kwarg — raises TypeError.
|
||||
Remove ddof from both rolling(..., ddof=N) and rolling(...).std(ddof=N).
|
||||
"""
|
||||
fixed_code = code
|
||||
|
||||
# Form 1: ddof inside rolling() — .rolling(window=N, min_periods=M, ddof=K)
|
||||
def _strip_ddof_from_rolling(m):
|
||||
inner = re.sub(r',?\s*ddof\s*=\s*\d+', '', m.group(1))
|
||||
inner = inner.strip(', ')
|
||||
self.fixes_applied.append("rolling_ddof: removed ddof from rolling()")
|
||||
return f'.rolling({inner})'
|
||||
|
||||
fixed_code = re.sub(r'\.rolling\(([^)]*ddof\s*=\s*\d+[^)]*)\)', _strip_ddof_from_rolling, fixed_code)
|
||||
|
||||
# Form 2: ddof inside .std() / .var() — .std(ddof=N)
|
||||
if re.search(r'\.(std|var)\([^)]*ddof\s*=\s*\d+', fixed_code):
|
||||
fixed_code = re.sub(r'\.(std|var)\([^)]*ddof\s*=\s*\d+[^)]*\)', r'.\1()', fixed_code)
|
||||
self.fixes_applied.append("rolling_ddof: removed ddof from std()/var()")
|
||||
|
||||
return fixed_code
|
||||
|
||||
def _fix_min_periods(self, code: str) -> str:
|
||||
"""
|
||||
Fix: Ensure min_periods matches window size in rolling calculations.
|
||||
@@ -325,6 +681,45 @@ 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()"
|
||||
)
|
||||
|
||||
# === 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*\)"
|
||||
|
||||
@@ -161,8 +161,7 @@ class FactorFBWorkspace(FBWorkspace):
|
||||
|
||||
try:
|
||||
subprocess.check_output(
|
||||
f"{FACTOR_COSTEER_SETTINGS.python_bin} {execution_code_path}",
|
||||
shell=True,
|
||||
[FACTOR_COSTEER_SETTINGS.python_bin, str(execution_code_path)],
|
||||
cwd=self.workspace_path,
|
||||
stderr=subprocess.STDOUT,
|
||||
timeout=FACTOR_COSTEER_SETTINGS.file_based_execution_timeout,
|
||||
|
||||
@@ -53,7 +53,7 @@ evolving_strategy_factor_implementation_v1_system: |-
|
||||
- ALWAYS use `min_periods=N` where N equals the window size in rolling calculations (e.g., `.rolling(20, min_periods=20)`)
|
||||
- ALWAYS handle infinite values after division: `.replace([np.inf, -np.inf], np.nan)` before saving results
|
||||
- ALWAYS use `groupby(level=1)` or `groupby('instrument')` before rolling operations on MultiIndex dataframes
|
||||
- Process the COMPLETE date range (2020-2026), do NOT filter by date
|
||||
- Process the COMPLETE date range available in the HDF5 file (do NOT filter by date — the file may contain 2024 debug data or full 2020-2026 data)
|
||||
- Use `groupby().transform()` instead of `groupby().apply()` for single-column assignments
|
||||
|
||||
Notice that you should not add any other text before or after the json format.
|
||||
|
||||
@@ -6,6 +6,7 @@ Two-step validation:
|
||||
2. Micro-batch testing - Runtime validation with small dataset
|
||||
"""
|
||||
|
||||
import ast
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
@@ -229,7 +230,7 @@ class LLMConfigValidator:
|
||||
final_metrics = re.search(r"\{'train_runtime':[^}]+\}", stdout)
|
||||
if final_metrics:
|
||||
try:
|
||||
metrics = eval(final_metrics.group(0)) # Safe: only numbers and strings
|
||||
metrics = ast.literal_eval(final_metrics.group(0))
|
||||
result["final_metrics"] = {
|
||||
"train_loss": metrics.get("train_loss"),
|
||||
"train_runtime": metrics.get("train_runtime"),
|
||||
|
||||
@@ -123,8 +123,8 @@ model_cls = AntiSymmetricConv
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
node_features = torch.load("node_features.pt")
|
||||
edge_index = torch.load("edge_index.pt")
|
||||
node_features = torch.load("node_features.pt", weights_only=True)
|
||||
edge_index = torch.load("edge_index.pt", weights_only=True)
|
||||
|
||||
# Model instantiation and forward pass
|
||||
model = AntiSymmetricConv(in_channels=node_features.size(-1))
|
||||
|
||||
@@ -78,8 +78,8 @@ model_cls = DirGNNConv
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
node_features = torch.load("node_features.pt")
|
||||
edge_index = torch.load("edge_index.pt")
|
||||
node_features = torch.load("node_features.pt", weights_only=True)
|
||||
edge_index = torch.load("edge_index.pt", weights_only=True)
|
||||
|
||||
# Model instantiation and forward pass
|
||||
model = DirGNNConv(MessagePassing())
|
||||
|
||||
@@ -187,8 +187,8 @@ model_cls = GPSConv
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
node_features = torch.load("node_features.pt")
|
||||
edge_index = torch.load("edge_index.pt")
|
||||
node_features = torch.load("node_features.pt", weights_only=True)
|
||||
edge_index = torch.load("edge_index.pt", weights_only=True)
|
||||
|
||||
# Model instantiation and forward pass
|
||||
model = GPSConv(channels=node_features.size(-1), conv=MessagePassing())
|
||||
|
||||
@@ -170,8 +170,8 @@ class LINKX(torch.nn.Module):
|
||||
model_cls = LINKX
|
||||
|
||||
if __name__ == "__main__":
|
||||
node_features = torch.load("node_features.pt")
|
||||
edge_index = torch.load("edge_index.pt")
|
||||
node_features = torch.load("node_features.pt", weights_only=True)
|
||||
edge_index = torch.load("edge_index.pt", weights_only=True)
|
||||
|
||||
# Model instantiation and forward pass
|
||||
model = LINKX(
|
||||
|
||||
@@ -102,8 +102,8 @@ class PMLP(torch.nn.Module):
|
||||
model_cls = PMLP
|
||||
|
||||
if __name__ == "__main__":
|
||||
node_features = torch.load("node_features.pt")
|
||||
edge_index = torch.load("edge_index.pt")
|
||||
node_features = torch.load("node_features.pt", weights_only=True)
|
||||
edge_index = torch.load("edge_index.pt", weights_only=True)
|
||||
|
||||
# Model instantiation and forward pass
|
||||
model = PMLP(
|
||||
|
||||
@@ -1180,8 +1180,8 @@ model_cls = ViSNet
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
node_features = torch.load("node_features.pt")
|
||||
edge_index = torch.load("edge_index.pt")
|
||||
node_features = torch.load("node_features.pt", weights_only=True)
|
||||
edge_index = torch.load("edge_index.pt", weights_only=True)
|
||||
|
||||
# Model instantiation and forward pass
|
||||
model = ViSNet()
|
||||
|
||||
@@ -58,10 +58,12 @@ class ModelCodeEvaluator(CoSTEEREvaluator):
|
||||
model_execution_feedback: str = "",
|
||||
model_value_feedback: str = "",
|
||||
):
|
||||
assert isinstance(target_task, ModelTask)
|
||||
assert isinstance(implementation, ModelFBWorkspace)
|
||||
if gt_implementation is not None:
|
||||
assert isinstance(gt_implementation, ModelFBWorkspace)
|
||||
if not isinstance(target_task, ModelTask):
|
||||
raise TypeError("target_task must be of type ModelTask")
|
||||
if not isinstance(implementation, ModelFBWorkspace):
|
||||
raise TypeError("implementation must be of type ModelFBWorkspace")
|
||||
if gt_implementation is not None and not isinstance(gt_implementation, ModelFBWorkspace):
|
||||
raise TypeError("gt_implementation must be of type ModelFBWorkspace")
|
||||
|
||||
model_task_information = target_task.get_task_information()
|
||||
code = implementation.all_codes
|
||||
@@ -113,10 +115,12 @@ class ModelFinalEvaluator(CoSTEEREvaluator):
|
||||
model_value_feedback: str,
|
||||
model_code_feedback: str,
|
||||
):
|
||||
assert isinstance(target_task, ModelTask)
|
||||
assert isinstance(implementation, ModelFBWorkspace)
|
||||
if gt_implementation is not None:
|
||||
assert isinstance(gt_implementation, ModelFBWorkspace)
|
||||
if not isinstance(target_task, ModelTask):
|
||||
raise TypeError("target_task must be of type ModelTask")
|
||||
if not isinstance(implementation, ModelFBWorkspace):
|
||||
raise TypeError("implementation must be of type ModelFBWorkspace")
|
||||
if gt_implementation is not None and not isinstance(gt_implementation, ModelFBWorkspace):
|
||||
raise TypeError("gt_implementation must be of type ModelFBWorkspace")
|
||||
|
||||
system_prompt = T(".prompts:evaluator_final_feedback.system").r(
|
||||
scenario=(
|
||||
|
||||
@@ -41,7 +41,8 @@ class ModelCoSTEEREvaluator(CoSTEEREvaluator):
|
||||
final_feedback="This task has failed too many times, skip implementation.",
|
||||
final_decision=False,
|
||||
)
|
||||
assert isinstance(target_task, ModelTask)
|
||||
if not isinstance(target_task, ModelTask):
|
||||
raise TypeError(f"Expected ModelTask, got {type(target_task)}")
|
||||
|
||||
# NOTE: Use fixed input to test the model to avoid randomness
|
||||
batch_size = 8
|
||||
@@ -50,7 +51,8 @@ class ModelCoSTEEREvaluator(CoSTEEREvaluator):
|
||||
input_value = 0.4
|
||||
param_init_value = 0.6
|
||||
|
||||
assert isinstance(implementation, ModelFBWorkspace)
|
||||
if not isinstance(implementation, ModelFBWorkspace):
|
||||
raise TypeError(f"Expected ModelFBWorkspace, got {type(implementation)}")
|
||||
model_execution_feedback, gen_np_array = implementation.execute(
|
||||
batch_size=batch_size,
|
||||
num_features=num_features,
|
||||
@@ -59,7 +61,8 @@ class ModelCoSTEEREvaluator(CoSTEEREvaluator):
|
||||
param_init_value=param_init_value,
|
||||
)
|
||||
if gt_implementation is not None:
|
||||
assert isinstance(gt_implementation, ModelFBWorkspace)
|
||||
if not isinstance(gt_implementation, ModelFBWorkspace):
|
||||
raise TypeError(f"Expected ModelFBWorkspace, got {type(gt_implementation)}")
|
||||
_, gt_np_array = gt_implementation.execute(
|
||||
batch_size=batch_size,
|
||||
num_features=num_features,
|
||||
|
||||
@@ -125,8 +125,8 @@ class AntiSymmetricConv(torch.nn.Module):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
node_features = torch.load("node_features.pt")
|
||||
edge_index = torch.load("edge_index.pt")
|
||||
node_features = torch.load("node_features.pt", weights_only=True)
|
||||
edge_index = torch.load("edge_index.pt", weights_only=True)
|
||||
|
||||
# Model instantiation and forward pass
|
||||
model = AntiSymmetricConv(in_channels=node_features.size(-1))
|
||||
|
||||
@@ -605,16 +605,15 @@ class OptunaOptimizer:
|
||||
synthetic_close = (1 + combined_ret).cumprod() * 100.0
|
||||
|
||||
from rdagent.components.backtesting.vbt_backtest import (
|
||||
backtest_signal,
|
||||
backtest_signal_ftmo,
|
||||
DEFAULT_TXN_COST_BPS,
|
||||
)
|
||||
import os as _os
|
||||
|
||||
bt = backtest_signal(
|
||||
bt = backtest_signal_ftmo(
|
||||
close=synthetic_close,
|
||||
signal=signal,
|
||||
txn_cost_bps=float(_os.getenv("TXN_COST_BPS", DEFAULT_TXN_COST_BPS)),
|
||||
freq="1min",
|
||||
)
|
||||
if bt.get("status") != "success":
|
||||
return self._default_metrics()
|
||||
|
||||
@@ -854,7 +854,7 @@ signal = signal.rolling(window=3, min_periods=1).mean().round().astype(int)
|
||||
# Delegate all metric computation to the single source of truth.
|
||||
# Same formulas as every other backtest path in the repo.
|
||||
from rdagent.components.backtesting.vbt_backtest import (
|
||||
backtest_signal,
|
||||
backtest_signal_ftmo,
|
||||
DEFAULT_TXN_COST_BPS,
|
||||
)
|
||||
|
||||
@@ -868,11 +868,10 @@ signal = signal.rolling(window=3, min_periods=1).mean().round().astype(int)
|
||||
close_for_bt = close.reindex(signal.index).ffill()
|
||||
|
||||
txn_cost_bps = float(os.getenv("TXN_COST_BPS", DEFAULT_TXN_COST_BPS))
|
||||
bt = backtest_signal(
|
||||
bt = backtest_signal_ftmo(
|
||||
close=close_for_bt,
|
||||
signal=signal,
|
||||
txn_cost_bps=txn_cost_bps,
|
||||
freq="1min",
|
||||
)
|
||||
|
||||
if bt.get("status") != "success":
|
||||
@@ -1265,7 +1264,7 @@ signal = signal.rolling(window=3, min_periods=1).mean().round().astype(int)
|
||||
signal = local_vars["signal"]
|
||||
|
||||
from rdagent.components.backtesting.vbt_backtest import (
|
||||
backtest_signal,
|
||||
backtest_signal_ftmo,
|
||||
DEFAULT_TXN_COST_BPS,
|
||||
)
|
||||
|
||||
@@ -1273,11 +1272,10 @@ signal = signal.rolling(window=3, min_periods=1).mean().round().astype(int)
|
||||
if close_for_bt is None:
|
||||
return {"sharpe_ratio": float('-inf'), "status": "rejected"}
|
||||
|
||||
bt = backtest_signal(
|
||||
bt = backtest_signal_ftmo(
|
||||
close=close_for_bt,
|
||||
signal=signal,
|
||||
txn_cost_bps=float(os.getenv("TXN_COST_BPS", DEFAULT_TXN_COST_BPS)),
|
||||
freq="1min",
|
||||
)
|
||||
if bt.get("status") != "success":
|
||||
return {"sharpe_ratio": float('-inf'), "status": "rejected"}
|
||||
|
||||
@@ -85,13 +85,16 @@ def load_and_process_one_pdf_by_azure_document_intelligence(
|
||||
|
||||
|
||||
def load_and_process_pdfs_by_azure_document_intelligence(path: Path) -> dict[str, str]:
|
||||
assert RD_AGENT_SETTINGS.azure_document_intelligence_key is not None
|
||||
assert RD_AGENT_SETTINGS.azure_document_intelligence_endpoint is not None
|
||||
if RD_AGENT_SETTINGS.azure_document_intelligence_key is None:
|
||||
raise AssertionError("azure_document_intelligence_key must be set")
|
||||
if RD_AGENT_SETTINGS.azure_document_intelligence_endpoint is None:
|
||||
raise AssertionError("azure_document_intelligence_endpoint must be set")
|
||||
|
||||
content_dict = {}
|
||||
ab_path = path.resolve()
|
||||
if ab_path.is_file():
|
||||
assert ".pdf" in ab_path.suffixes, "The file must be a PDF file."
|
||||
if ".pdf" not in ab_path.suffixes:
|
||||
raise ValueError("The file must be a PDF file.")
|
||||
proc = load_and_process_one_pdf_by_azure_document_intelligence
|
||||
content_dict[str(ab_path)] = proc(
|
||||
ab_path,
|
||||
|
||||
@@ -24,7 +24,8 @@ class UndirectedNode(Node):
|
||||
super().__init__(content, label, embedding)
|
||||
self.neighbors: set[UndirectedNode] = set()
|
||||
self.appendix = appendix # appendix stores any additional information
|
||||
assert isinstance(content, str), "content must be a string"
|
||||
if not isinstance(content, str):
|
||||
raise TypeError("content must be a string")
|
||||
|
||||
def add_neighbor(self, node: UndirectedNode) -> None:
|
||||
self.neighbors.add(node)
|
||||
@@ -96,7 +97,8 @@ class Graph(KnowledgeBase):
|
||||
APIBackend().create_embedding(input_content=contents[i : i + size]),
|
||||
)
|
||||
|
||||
assert len(nodes) == len(embeddings), "nodes' length must equals embeddings' length"
|
||||
if len(nodes) != len(embeddings):
|
||||
raise ValueError("nodes' length must equal embeddings' length")
|
||||
for node, embedding in zip(nodes, embeddings):
|
||||
node.embedding = embedding
|
||||
return nodes
|
||||
@@ -252,7 +254,8 @@ class UndirectedGraph(Graph):
|
||||
|
||||
"""
|
||||
min_nodes_count = 2
|
||||
assert len(nodes) >= min_nodes_count, "nodes length must >=2"
|
||||
if len(nodes) < min_nodes_count:
|
||||
raise ValueError("nodes length must >=2")
|
||||
intersection = None
|
||||
|
||||
for node in nodes:
|
||||
|
||||
@@ -87,7 +87,8 @@ class ModelWsLoader(WsLoader[ModelTask, ModelFBWorkspace]):
|
||||
self.path = Path(path)
|
||||
|
||||
def load(self, task: ModelTask) -> ModelFBWorkspace:
|
||||
assert task.name is not None
|
||||
if task.name is None:
|
||||
raise AssertionError("task.name should not be None")
|
||||
mti = ModelFBWorkspace(task)
|
||||
mti.prepare()
|
||||
with open(self.path / f"{task.name}.py", "r") as f:
|
||||
|
||||
@@ -4,6 +4,7 @@ import functools
|
||||
import importlib
|
||||
import json
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
from collections.abc import Callable
|
||||
@@ -208,3 +209,13 @@ def cache_with_pickle(hash_func: Callable, post_process_func: Callable | None =
|
||||
return cache_wrapper
|
||||
|
||||
return cache_decorator
|
||||
|
||||
|
||||
def safe_resolve_path(user_path: Path, safe_root: Path | None = None) -> Path:
|
||||
if safe_root is not None:
|
||||
root_real = os.path.realpath(str(safe_root.expanduser()))
|
||||
path_real = os.path.realpath(str(user_path.expanduser())) # nosec B614 — validated against safe_root below
|
||||
if not (path_real == root_real or path_real.startswith(root_real + os.sep)):
|
||||
raise ValueError(f"Path {user_path} resolves to {path_real}, outside allowed root {safe_root}")
|
||||
return Path(path_real)
|
||||
return user_path.expanduser().resolve()
|
||||
|
||||
+22
-15
@@ -27,6 +27,7 @@ Usage:
|
||||
from __future__ import annotations
|
||||
|
||||
import json as _json
|
||||
import logging
|
||||
import sys
|
||||
import threading
|
||||
from contextlib import contextmanager
|
||||
@@ -36,21 +37,24 @@ from typing import Any
|
||||
|
||||
from loguru import logger as _root
|
||||
|
||||
# ── paths ─────────────────────────────────────────────────────────────────────
|
||||
# ── paths ─────────────────────────────────────────────────────────────────────────────────
|
||||
LOGS_ROOT: Path = Path(__file__).parent.parent.parent / "logs"
|
||||
|
||||
# ── format ────────────────────────────────────────────────────────────────────
|
||||
# ── format ────────────────────────────────────────────────────────────────────────────────
|
||||
_FILE_FMT = (
|
||||
"{time:YYYY-MM-DD HH:mm:ss.SSS} | {level: <8} | {extra[cmd]: <18} | {message}"
|
||||
)
|
||||
|
||||
# ── internal state ─────────────────────────────────────────────────────────────
|
||||
# ── internal state ─────────────────────────────────────────────────────────────────────────────
|
||||
_registered: set[str] = set() # command keys that already have a file sink
|
||||
_all_added: bool = False # whether the combined all.log sink is active
|
||||
_llm_log_lock = threading.Lock() # guards concurrent writes to llm_calls.jsonl
|
||||
|
||||
# Maximum characters stored per field in llm_calls.jsonl to prevent GB-scale files.
|
||||
_LLM_CALL_MAX_CHARS = 500
|
||||
|
||||
# ── helpers ───────────────────────────────────────────────────────────────────
|
||||
|
||||
# ── helpers ────────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
def _today_dir() -> Path:
|
||||
d = LOGS_ROOT / datetime.now().strftime("%Y-%m-%d")
|
||||
@@ -79,7 +83,7 @@ def _banner(log, title: str, meta: dict[str, Any]) -> None:
|
||||
log.info(sep)
|
||||
|
||||
|
||||
# ── public API ────────────────────────────────────────────────────────────────
|
||||
# ── public API ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
def log_llm_call(
|
||||
system: str | None,
|
||||
@@ -88,16 +92,19 @@ def log_llm_call(
|
||||
start_time: Any = None,
|
||||
end_time: Any = None,
|
||||
) -> None:
|
||||
"""Append one complete LLM call to logs/YYYY-MM-DD/llm_calls.jsonl.
|
||||
"""Append one LLM call summary to logs/YYYY-MM-DD/llm_calls.jsonl.
|
||||
|
||||
Prompt/response content is capped at _LLM_CALL_MAX_CHARS to prevent
|
||||
GB-scale log files from long-running loops.
|
||||
|
||||
Each line is a self-contained JSON object so the file is grep/jq-friendly:
|
||||
jq 'select(.duration_ms > 5000)' logs/2026-04-17/llm_calls.jsonl
|
||||
"""
|
||||
entry: dict[str, Any] = {
|
||||
"ts": datetime.now().isoformat(timespec="milliseconds"),
|
||||
"system": system or "",
|
||||
"user": user,
|
||||
"response": response,
|
||||
"system": (system or "")[:_LLM_CALL_MAX_CHARS],
|
||||
"user": user[:_LLM_CALL_MAX_CHARS],
|
||||
"response": response[:_LLM_CALL_MAX_CHARS],
|
||||
}
|
||||
if start_time is not None and end_time is not None:
|
||||
try:
|
||||
@@ -130,13 +137,13 @@ def setup(command: str, **context: Any):
|
||||
key = command.lower()
|
||||
|
||||
if key not in _registered:
|
||||
# Per-command rotating file
|
||||
_root.add(
|
||||
str(log_dir / f"{key}.log"),
|
||||
format=_FILE_FMT,
|
||||
filter=lambda r, k=key: r["extra"].get("cmd", "").lower() == k,
|
||||
rotation="00:00", # new file at midnight
|
||||
retention="30 days",
|
||||
rotation="50 MB",
|
||||
compression="gz",
|
||||
retention="7 days",
|
||||
encoding="utf-8",
|
||||
enqueue=True,
|
||||
backtrace=False,
|
||||
@@ -145,13 +152,13 @@ def setup(command: str, **context: Any):
|
||||
_registered.add(key)
|
||||
|
||||
if not _all_added:
|
||||
# Combined log — all commands
|
||||
_root.add(
|
||||
str(log_dir / "all.log"),
|
||||
format=_FILE_FMT,
|
||||
filter=lambda r: "cmd" in r["extra"],
|
||||
rotation="00:00",
|
||||
retention="60 days",
|
||||
rotation="100 MB",
|
||||
compression="gz",
|
||||
retention="7 days",
|
||||
encoding="utf-8",
|
||||
enqueue=True,
|
||||
backtrace=False,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,306 @@
|
||||
import argparse
|
||||
import json
|
||||
import pickle # nosec
|
||||
import re
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import streamlit as st
|
||||
from streamlit import session_state
|
||||
|
||||
from rdagent.log.ui.conf import UI_SETTING
|
||||
from rdagent.log.utils import extract_evoid, extract_loopid_func_name
|
||||
|
||||
st.set_page_config(layout="wide", page_title="debug_llm", page_icon="🎓", initial_sidebar_state="expanded")
|
||||
|
||||
# 获取 log_path 参数
|
||||
parser = argparse.ArgumentParser(description="RD-Agent Streamlit App")
|
||||
parser.add_argument("--log_dir", type=str, help="Path to the log directory")
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
def get_folders_sorted(log_path):
|
||||
"""缓存并返回排序后的文件夹列表,并加入进度打印"""
|
||||
with st.spinner("正在加载文件夹列表..."):
|
||||
folders = sorted(
|
||||
(folder for folder in log_path.iterdir() if folder.is_dir() and list(folder.iterdir())),
|
||||
key=lambda folder: folder.stat().st_mtime,
|
||||
reverse=True,
|
||||
)
|
||||
st.write(f"找到 {len(folders)} 个文件夹")
|
||||
return [folder.name for folder in folders]
|
||||
|
||||
|
||||
if UI_SETTING.enable_cache:
|
||||
get_folders_sorted = st.cache_data(get_folders_sorted)
|
||||
|
||||
|
||||
# 设置主日志路径
|
||||
main_log_path = Path(args.log_dir) if args.log_dir else Path("./log")
|
||||
if not main_log_path.exists():
|
||||
st.error(f"Log dir {main_log_path} does not exist!")
|
||||
st.stop()
|
||||
|
||||
if "data" not in session_state:
|
||||
session_state.data = []
|
||||
if "log_path" not in session_state:
|
||||
session_state.log_path = None
|
||||
|
||||
tlist = []
|
||||
|
||||
|
||||
def load_data():
|
||||
"""加载数据到 session_state 并显示进度"""
|
||||
log_file = main_log_path / session_state.log_path / "debug_llm.pkl"
|
||||
try:
|
||||
with st.spinner(f"正在加载数据文件 {log_file}..."):
|
||||
start_time = time.time()
|
||||
with open(log_file, "rb") as f:
|
||||
session_state.data = pickle.load(f, encoding="utf-8") # nosec
|
||||
st.success(f"数据加载完成!耗时 {time.time() - start_time:.2f} 秒")
|
||||
st.session_state["current_loop"] = 1
|
||||
except Exception as e:
|
||||
session_state.data = [{"error": str(e)}]
|
||||
st.error(f"加载数据失败: {e}")
|
||||
|
||||
|
||||
# UI - Sidebar
|
||||
with st.sidebar:
|
||||
st.markdown(":blue[**Log Path**]")
|
||||
manually = st.toggle("Manual Input")
|
||||
if manually:
|
||||
st.text_input("log path", key="log_path", label_visibility="collapsed")
|
||||
else:
|
||||
folders = get_folders_sorted(main_log_path)
|
||||
st.selectbox(f"**Select from {main_log_path.absolute()}**", folders, key="log_path") # nosec B608 — not SQL, Bandit false positive on "Select" in UI label
|
||||
|
||||
if st.button("Refresh Data"):
|
||||
load_data()
|
||||
st.rerun()
|
||||
|
||||
|
||||
# Helper functions
|
||||
def show_text(text, lang=None):
|
||||
"""显示文本代码块"""
|
||||
if lang:
|
||||
st.code(text, language=lang, wrap_lines=True)
|
||||
elif "\n" in text:
|
||||
st.code(text, language="python", wrap_lines=True)
|
||||
else:
|
||||
st.code(text, language="html", wrap_lines=True)
|
||||
|
||||
|
||||
def highlight_prompts_uri(uri):
|
||||
"""高亮 URI 的格式"""
|
||||
parts = uri.split(":")
|
||||
return f"**{parts[0]}:**:green[**{parts[1]}**]"
|
||||
|
||||
|
||||
# Display Data
|
||||
progress_text = st.empty()
|
||||
progress_bar = st.progress(0)
|
||||
|
||||
# 每页展示一个 Loop
|
||||
LOOPS_PER_PAGE = 1
|
||||
|
||||
# 获取所有的 Loop ID
|
||||
loop_groups = {}
|
||||
for i, d in enumerate(session_state.data):
|
||||
tag = d["tag"]
|
||||
loop_id, _ = extract_loopid_func_name(tag)
|
||||
if loop_id:
|
||||
if loop_id not in loop_groups:
|
||||
loop_groups[loop_id] = []
|
||||
loop_groups[loop_id].append(d)
|
||||
|
||||
# 按 Loop ID 排序
|
||||
sorted_loop_ids = sorted(loop_groups.keys(), key=int) # 假设 Loop ID 是数字
|
||||
total_loops = len(sorted_loop_ids)
|
||||
total_pages = total_loops # 每页展示一个 Loop
|
||||
|
||||
|
||||
# simple display
|
||||
# FIXME: Delete this simple UI if trace have tag(evo_id & loop_id)
|
||||
# with st.sidebar:
|
||||
# start = int(st.text_input("start", 0))
|
||||
# end = int(st.text_input("end", 100))
|
||||
# for m in session_state.data[start:end]:
|
||||
# if "tpl" in m["tag"]:
|
||||
# obj = m["obj"]
|
||||
# uri = obj["uri"]
|
||||
# tpl = obj["template"]
|
||||
# cxt = obj["context"]
|
||||
# rd = obj["rendered"]
|
||||
# with st.expander(highlight_prompts_uri(uri), expanded=False, icon="⚙️"):
|
||||
# t1, t2, t3 = st.tabs([":green[**Rendered**]", ":blue[**Template**]", ":orange[**Context**]"])
|
||||
# with t1:
|
||||
# show_text(rd)
|
||||
# with t2:
|
||||
# show_text(tpl, lang="django")
|
||||
# with t3:
|
||||
# st.json(cxt)
|
||||
# if "llm" in m["tag"]:
|
||||
# obj = m["obj"]
|
||||
# system = obj.get("system", None)
|
||||
# user = obj["user"]
|
||||
# resp = obj["resp"]
|
||||
# with st.expander(f"**LLM**", expanded=False, icon="🤖"):
|
||||
# t1, t2, t3 = st.tabs([":green[**Response**]", ":blue[**User**]", ":orange[**System**]"])
|
||||
# with t1:
|
||||
# try:
|
||||
# rdict = json.loads(resp)
|
||||
# if "code" in rdict:
|
||||
# code = rdict["code"]
|
||||
# st.markdown(":red[**Code in response dict:**]")
|
||||
# st.code(code, language="python", wrap_lines=True, line_numbers=True)
|
||||
# rdict.pop("code")
|
||||
# elif "spec" in rdict:
|
||||
# spec = rdict["spec"]
|
||||
# st.markdown(":red[**Spec in response dict:**]")
|
||||
# st.markdown(spec)
|
||||
# rdict.pop("spec")
|
||||
# else:
|
||||
# # show model codes
|
||||
# showed_keys = []
|
||||
# for k, v in rdict.items():
|
||||
# if k.startswith("model_") and k.endswith(".py"):
|
||||
# st.markdown(f":red[**{k}**]")
|
||||
# st.code(v, language="python", wrap_lines=True, line_numbers=True)
|
||||
# showed_keys.append(k)
|
||||
# for k in showed_keys:
|
||||
# rdict.pop(k)
|
||||
# st.write(":red[**Other parts (except for the code or spec) in response dict:**]")
|
||||
# st.json(rdict)
|
||||
# except:
|
||||
# st.json(resp)
|
||||
# with t2:
|
||||
# show_text(user)
|
||||
# with t3:
|
||||
# show_text(system or "No system prompt available")
|
||||
|
||||
|
||||
if total_pages:
|
||||
# 初始化 current_loop
|
||||
if "current_loop" not in st.session_state:
|
||||
st.session_state["current_loop"] = 1
|
||||
|
||||
# Loop 导航按钮
|
||||
col1, col2, col3, col4, col5 = st.sidebar.columns([1.2, 1, 2, 1, 1.2])
|
||||
|
||||
with col1:
|
||||
if st.button("|<"): # 首页
|
||||
st.session_state["current_loop"] = 1
|
||||
with col2:
|
||||
if st.button("<") and st.session_state["current_loop"] > 1: # 上一页
|
||||
st.session_state["current_loop"] -= 1
|
||||
with col3:
|
||||
# 下拉列表显示所有 Loop
|
||||
st.session_state["current_loop"] = st.selectbox(
|
||||
"选择 Loop",
|
||||
options=list(range(1, total_loops + 1)),
|
||||
index=st.session_state["current_loop"] - 1, # 默认选中当前 Loop
|
||||
label_visibility="collapsed", # 隐藏标签
|
||||
)
|
||||
with col4:
|
||||
if st.button("\>") and st.session_state["current_loop"] < total_loops: # 下一页
|
||||
st.session_state["current_loop"] += 1
|
||||
with col5:
|
||||
if st.button("\>|"): # 最后一页
|
||||
st.session_state["current_loop"] = total_loops
|
||||
|
||||
# 获取当前 Loop
|
||||
current_loop = st.session_state["current_loop"]
|
||||
|
||||
# 渲染当前 Loop 数据
|
||||
loop_id = sorted_loop_ids[current_loop - 1]
|
||||
progress_text = st.empty()
|
||||
progress_text.text(f"正在处理 Loop {loop_id}...")
|
||||
progress_bar.progress(current_loop / total_loops, text=f"Loop :green[**{current_loop}**] / {total_loops}")
|
||||
|
||||
# 渲染 Loop Header
|
||||
loop_anchor = f"Loop_{loop_id}"
|
||||
if loop_anchor not in tlist:
|
||||
tlist.append(loop_anchor)
|
||||
st.header(loop_anchor, anchor=loop_anchor, divider="blue")
|
||||
|
||||
# 渲染当前 Loop 的所有数据
|
||||
loop_data = loop_groups[loop_id]
|
||||
for d in loop_data:
|
||||
tag = d["tag"]
|
||||
obj = d["obj"]
|
||||
_, func_name = extract_loopid_func_name(tag)
|
||||
evo_id = extract_evoid(tag)
|
||||
|
||||
func_anchor = f"loop_{loop_id}.{func_name}"
|
||||
if func_anchor not in tlist:
|
||||
tlist.append(func_anchor)
|
||||
st.header(f"in *{func_name}*", anchor=func_anchor, divider="green")
|
||||
|
||||
evo_anchor = f"loop_{loop_id}.evo_step_{evo_id}"
|
||||
if evo_id and evo_anchor not in tlist:
|
||||
tlist.append(evo_anchor)
|
||||
st.subheader(f"evo_step_{evo_id}", anchor=evo_anchor, divider="orange")
|
||||
|
||||
# 根据 tag 渲染内容
|
||||
if "debug_exp_gen" in tag:
|
||||
with st.expander(
|
||||
f"Exp in :violet[**{obj.experiment_workspace.workspace_path}**]", expanded=False, icon="🧩"
|
||||
):
|
||||
st.write(obj)
|
||||
elif "debug_tpl" in tag:
|
||||
uri = obj["uri"]
|
||||
tpl = obj["template"]
|
||||
cxt = obj["context"]
|
||||
rd = obj["rendered"]
|
||||
with st.expander(highlight_prompts_uri(uri), expanded=False, icon="⚙️"):
|
||||
t1, t2, t3 = st.tabs([":green[**Rendered**]", ":blue[**Template**]", ":orange[**Context**]"])
|
||||
with t1:
|
||||
show_text(rd)
|
||||
with t2:
|
||||
show_text(tpl, lang="django")
|
||||
with t3:
|
||||
st.json(cxt)
|
||||
elif "debug_llm" in tag:
|
||||
system = obj.get("system", None)
|
||||
user = obj["user"]
|
||||
resp = obj["resp"]
|
||||
with st.expander(f"**LLM**", expanded=False, icon="🤖"):
|
||||
t1, t2, t3 = st.tabs([":green[**Response**]", ":blue[**User**]", ":orange[**System**]"])
|
||||
with t1:
|
||||
try:
|
||||
rdict = json.loads(resp)
|
||||
if "code" in rdict:
|
||||
code = rdict["code"]
|
||||
st.markdown(":red[**Code in response dict:**]")
|
||||
st.code(code, language="python", wrap_lines=True, line_numbers=True)
|
||||
rdict.pop("code")
|
||||
elif "spec" in rdict:
|
||||
spec = rdict["spec"]
|
||||
st.markdown(":red[**Spec in response dict:**]")
|
||||
st.markdown(spec)
|
||||
rdict.pop("spec")
|
||||
else:
|
||||
# show model codes
|
||||
showed_keys = []
|
||||
for k, v in rdict.items():
|
||||
if k.startswith("model_") and k.endswith(".py"):
|
||||
st.markdown(f":red[**{k}**]")
|
||||
st.code(v, language="python", wrap_lines=True, line_numbers=True)
|
||||
showed_keys.append(k)
|
||||
for k in showed_keys:
|
||||
rdict.pop(k)
|
||||
st.write(":red[**Other parts (except for the code or spec) in response dict:**]")
|
||||
st.json(rdict)
|
||||
except:
|
||||
st.json(resp)
|
||||
with t2:
|
||||
show_text(user)
|
||||
with t3:
|
||||
show_text(system or "No system prompt available")
|
||||
|
||||
progress_text.text("当前 Loop 数据处理完成!")
|
||||
|
||||
# Sidebar TOC
|
||||
with st.sidebar:
|
||||
toc = "\n".join([f"- [{t}](#{t})" if t.startswith("L") else f" - [{t.split('.')[1]}](#{t})" for t in tlist])
|
||||
st.markdown(toc, unsafe_allow_html=True)
|
||||
@@ -541,7 +541,8 @@ class APIBackend(ABC):
|
||||
**kwargs,
|
||||
) -> str | list[list[float]]:
|
||||
"""This function to share operation between embedding and chat completion"""
|
||||
assert not (chat_completion and embedding), "chat_completion and embedding cannot be True at the same time"
|
||||
if chat_completion and embedding:
|
||||
raise ValueError("chat_completion and embedding cannot be True at the same time")
|
||||
max_retry = LLM_SETTINGS.max_retry if LLM_SETTINGS.max_retry is not None else max_retry
|
||||
timeout_count = 0
|
||||
violation_count = 0
|
||||
@@ -720,7 +721,13 @@ class APIBackend(ABC):
|
||||
|
||||
if finish_reason is None or finish_reason != "length":
|
||||
break # we get a full response now.
|
||||
new_messages.append({"role": "assistant", "content": response})
|
||||
# Merge into the previous assistant message if there already is one at the end.
|
||||
# Appending a second consecutive assistant message causes llama-server to return 400
|
||||
# ("Cannot have 2 or more assistant messages at the end of the list").
|
||||
if new_messages and new_messages[-1]["role"] == "assistant":
|
||||
new_messages[-1]["content"] += response
|
||||
else:
|
||||
new_messages.append({"role": "assistant", "content": response})
|
||||
else:
|
||||
raise RuntimeError(f"Failed to continue the conversation after {try_n} retries.")
|
||||
|
||||
|
||||
@@ -36,16 +36,18 @@ def get_agent_model() -> OpenAIChatModel:
|
||||
|
||||
"""
|
||||
backend = APIBackend()
|
||||
assert isinstance(backend, LiteLLMAPIBackend), "Only LiteLLMAPIBackend is supported"
|
||||
if not isinstance(backend, LiteLLMAPIBackend):
|
||||
raise TypeError("Only LiteLLMAPIBackend is supported")
|
||||
|
||||
compl_kwargs = backend.get_complete_kwargs()
|
||||
|
||||
selected_model = compl_kwargs["model"]
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(selected_model)
|
||||
assert (
|
||||
custom_llm_provider in PROVIDER_TO_ENV_MAP
|
||||
), f"Provider {custom_llm_provider} not supported. Please add it into `PROVIDER_TO_ENV_MAP`"
|
||||
if custom_llm_provider not in PROVIDER_TO_ENV_MAP:
|
||||
raise ValueError(
|
||||
f"Provider {custom_llm_provider} not supported. Please add it into `PROVIDER_TO_ENV_MAP`"
|
||||
)
|
||||
prefix = PROVIDER_TO_ENV_MAP[custom_llm_provider]
|
||||
api_key = os.getenv(f"{prefix}_API_KEY", None)
|
||||
api_base = os.getenv(f"{prefix}_API_BASE", None)
|
||||
|
||||
@@ -268,7 +268,8 @@ class JsonReducer(DataReducer):
|
||||
parent[key] = sampled # type: ignore # parent 是 list,key 是 index, list.__setitem__(key, sampled)
|
||||
self.sampled_files.extend([self.extract_filename(i) for i in sampled])
|
||||
break
|
||||
assert len(self.sampled_files) > 0
|
||||
if len(self.sampled_files) <= 0:
|
||||
raise AssertionError("sampled_files must contain at least one file")
|
||||
return data
|
||||
|
||||
def _find_all_lists(
|
||||
|
||||
@@ -7,8 +7,10 @@ from sklearn.metrics import roc_auc_score
|
||||
def prepare_for_auroc_metric(submission: pd.DataFrame, answers: pd.DataFrame, id_col: str, target_col: str) -> dict:
|
||||
|
||||
# Answers checks
|
||||
assert id_col in answers.columns, f"answers dataframe should have an {id_col} column"
|
||||
assert target_col in answers.columns, f"answers dataframe should have a {target_col} column"
|
||||
if id_col not in answers.columns:
|
||||
raise InvalidSubmissionError(f"answers dataframe should have an {id_col} column")
|
||||
if target_col not in answers.columns:
|
||||
raise InvalidSubmissionError(f"answers dataframe should have a {target_col} column")
|
||||
|
||||
# Submission checks
|
||||
if id_col not in submission.columns:
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
from pathlib import Path
|
||||
|
||||
# Check if our submission file exists
|
||||
assert Path("submission.csv").exists(), "Error: submission.csv not found"
|
||||
if not Path("submission.csv").exists():
|
||||
raise FileNotFoundError("Error: submission.csv not found")
|
||||
|
||||
submission_lines = Path("submission.csv").read_text().splitlines()
|
||||
test_lines = Path("submission_test.csv").read_text().splitlines()
|
||||
|
||||
@@ -22,7 +22,8 @@ def prepare_for_metric(submission: pd.DataFrame, answers: pd.DataFrame) -> dict:
|
||||
if "price" not in submission.columns:
|
||||
raise InvalidSubmissionError("Submission DataFrame must contain 'price' columns.")
|
||||
|
||||
assert "price" in answers.columns, "Answers DataFrame must contain 'price' columns."
|
||||
if "price" not in answers.columns:
|
||||
raise InvalidSubmissionError("Answers DataFrame must contain 'price' columns.")
|
||||
|
||||
if len(submission) != len(answers):
|
||||
raise InvalidSubmissionError("Submission must be the same length as the answers.")
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
from pathlib import Path
|
||||
|
||||
# Check if our submission file exists
|
||||
assert Path("submission.csv").exists(), "Error: submission.csv not found"
|
||||
if not Path("submission.csv").exists():
|
||||
raise FileNotFoundError("Error: submission.csv not found")
|
||||
|
||||
submission_lines = Path("submission.csv").read_text().splitlines() # 自动生成的
|
||||
test_lines = Path("submission_test.csv").read_text().splitlines() # test.csv
|
||||
|
||||
+14
-11
@@ -56,14 +56,17 @@ sparse.save_npz(public / "test" / "X.npz", X_test)
|
||||
sparse.save_npz(public / "train" / "X.npz", X_train)
|
||||
df_train.to_csv(public / "train" / "ARF_12h.csv", index=False)
|
||||
|
||||
assert (
|
||||
X_train.shape[0] == df_train.shape[0]
|
||||
), f"Mismatch: X_train rows ({X_train.shape[0]}) != df_train rows ({df_train.shape[0]})"
|
||||
assert (
|
||||
X_test.shape[0] == df_test.shape[0]
|
||||
), f"Mismatch: X_test rows ({X_test.shape[0]}) != df_test rows ({df_test.shape[0]})"
|
||||
assert df_test.shape[1] == 2, "Public test set should have 2 columns"
|
||||
assert df_train.shape[1] == 3, "Public train set should have 3 columns"
|
||||
assert len(df_train) + len(df_test) == len(
|
||||
df_label
|
||||
), "Length of new_train and new_test should equal length of old_train"
|
||||
if X_train.shape[0] != df_train.shape[0]:
|
||||
raise ValueError(
|
||||
f"Mismatch: X_train rows ({X_train.shape[0]}) != df_train rows ({df_train.shape[0]})"
|
||||
)
|
||||
if X_test.shape[0] != df_test.shape[0]:
|
||||
raise ValueError(
|
||||
f"Mismatch: X_test rows ({X_test.shape[0]}) != df_test rows ({df_test.shape[0]})"
|
||||
)
|
||||
if df_test.shape[1] != 2:
|
||||
raise ValueError("Public test set should have 2 columns")
|
||||
if df_train.shape[1] != 3:
|
||||
raise ValueError("Public train set should have 3 columns")
|
||||
if len(df_train) + len(df_test) != len(df_label):
|
||||
raise ValueError("Length of new_train and new_test should equal length of old_train")
|
||||
|
||||
+6
-5
@@ -25,11 +25,12 @@ def prepare(raw: Path, public: Path, private: Path):
|
||||
new_test.to_csv(public / "test.csv", index=False)
|
||||
|
||||
# Checks
|
||||
assert new_test.shape[1] == 12, "Public test set should have 12 columns"
|
||||
assert new_train.shape[1] == 13, "Public train set should have 13 columns"
|
||||
assert len(new_train) + len(new_test) == len(
|
||||
old_train
|
||||
), "Length of new_train and new_test should equal length of old_train"
|
||||
if new_test.shape[1] != 12:
|
||||
raise AssertionError("Public test set should have 12 columns")
|
||||
if new_train.shape[1] != 13:
|
||||
raise AssertionError("Public train set should have 13 columns")
|
||||
if len(new_train) + len(new_test) != len(old_train):
|
||||
raise AssertionError("Length of new_train and new_test should equal length of old_train")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -320,7 +320,8 @@ class DataScienceRDLoop(RDLoop):
|
||||
# only clean current workspace without affecting other loops.
|
||||
for k in "direct_exp_gen", "coding", "running":
|
||||
if k in prev_out and prev_out[k] is not None:
|
||||
assert isinstance(prev_out[k], DSExperiment)
|
||||
if not isinstance(prev_out[k], DSExperiment):
|
||||
raise TypeError(f"prev_out[{k!r}] must be an instance of DSExperiment")
|
||||
clean_workspace(prev_out[k].experiment_workspace.workspace_path)
|
||||
|
||||
# Backup the workspace (only necessary files are included)
|
||||
|
||||
@@ -213,7 +213,8 @@ class DSTrace(Trace[DataScienceScen, KnowledgeBase]):
|
||||
self, component: COMPONENT, search_list: list[tuple[DSExperiment, ExperimentFeedback]] = []
|
||||
) -> bool:
|
||||
for exp, fb in search_list:
|
||||
assert isinstance(exp.hypothesis, DSHypothesis), "Hypothesis should be DSHypothesis (and not None)"
|
||||
if not isinstance(exp.hypothesis, DSHypothesis):
|
||||
raise TypeError("Hypothesis should be DSHypothesis (and not None)")
|
||||
if exp.hypothesis.component == component and fb:
|
||||
return True
|
||||
return False
|
||||
|
||||
@@ -182,7 +182,7 @@ class ExpGen2Hypothesis(DSProposalV2ExpGen):
|
||||
|
||||
success_fb_list = list(set(trace_fbs))
|
||||
logger.info(
|
||||
f"Merge Hypothesis: select {len(success_fb_list)} from {len(trace_fbs)} SOTA experiments found in {len(leaves)} traces"
|
||||
f"Merge Hypothesis: select {len(success_fb_list)} from {len(trace_fbs)} SOTA experiments found in {len(leaves)} traces" # nosec B608 — not SQL, Bandit false positive on "select" in log message
|
||||
)
|
||||
|
||||
if len(success_fb_list) > 0:
|
||||
@@ -377,7 +377,8 @@ class ExpGen2TraceAndMergeV2(ExpGen):
|
||||
if DS_RD_SETTING.enable_multi_version_exp_gen:
|
||||
exp_gen_version_list = DS_RD_SETTING.exp_gen_version_list.split(",")
|
||||
for version in exp_gen_version_list:
|
||||
assert version in ["v3", "v2", "v1"]
|
||||
if version not in ["v3", "v2", "v1"]:
|
||||
raise ValueError(f"version must be 'v1', 'v2', or 'v3', got {version!r}")
|
||||
|
||||
if len(trace.hist) == 0:
|
||||
# set the proposal version for the first sub-trace
|
||||
|
||||
@@ -339,7 +339,8 @@ class DSProposalV1ExpGen(ExpGen):
|
||||
eda_output = sota_exp.experiment_workspace.file_dict.get("EDA.md", None)
|
||||
scenario_desc = trace.scen.get_scenario_all_desc(eda_output=eda_output)
|
||||
|
||||
assert sota_exp is not None, "SOTA experiment is not provided."
|
||||
if sota_exp is None:
|
||||
raise ValueError("SOTA experiment is not provided.")
|
||||
last_exp = trace.last_exp()
|
||||
# exp_and_feedback = trace.hist[-1]
|
||||
# last_exp = exp_and_feedback[0]
|
||||
@@ -445,8 +446,10 @@ class DSProposalV1ExpGen(ExpGen):
|
||||
json_target_type=dict[str, dict[str, str | dict] | str],
|
||||
)
|
||||
)
|
||||
assert "hypothesis_proposal" in resp_dict, "Hypothesis proposal not provided."
|
||||
assert "task_design" in resp_dict, "Task design not provided."
|
||||
if "hypothesis_proposal" not in resp_dict:
|
||||
raise ValueError("Hypothesis proposal not provided.")
|
||||
if "task_design" not in resp_dict:
|
||||
raise ValueError("Task design not provided.")
|
||||
task_class = component_info["task_class"]
|
||||
hypothesis_proposal = resp_dict.get("hypothesis_proposal", {})
|
||||
hypothesis = DSHypothesis(
|
||||
@@ -1149,8 +1152,10 @@ You help users retrieve relevant knowledge from community discussions and public
|
||||
)
|
||||
|
||||
response_dict = json.loads(response)
|
||||
assert response_dict.get("component") in HypothesisComponent.__members__, f"Invalid component"
|
||||
assert response_dict.get("hypothesis") is not None, f"Invalid hypothesis"
|
||||
if response_dict.get("component") not in HypothesisComponent.__members__:
|
||||
raise ValueError(f"Invalid component: {response_dict.get('component')}")
|
||||
if response_dict.get("hypothesis") is None:
|
||||
raise ValueError("Invalid hypothesis")
|
||||
return response_dict
|
||||
|
||||
# END: for support llm-based hypothesis selection -----
|
||||
@@ -1253,7 +1258,8 @@ You help users retrieve relevant knowledge from community discussions and public
|
||||
description=task_desc,
|
||||
)
|
||||
|
||||
assert isinstance(task, PipelineTask), f"Task {task_name} is not a PipelineTask, got {type(task)}"
|
||||
if not isinstance(task, PipelineTask):
|
||||
raise TypeError(f"Task {task_name} is not a PipelineTask, got {type(task)}")
|
||||
# only for llm with response schema.(TODO: support for non-schema llm?)
|
||||
# If the LLM provides a "packages" field (list[str]), compute runtime environment now and cache it for subsequent prompts in later loops.
|
||||
if isinstance(task_dict, dict) and "packages" in task_dict and isinstance(task_dict["packages"], list):
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import ast
|
||||
import json
|
||||
import os
|
||||
import pickle
|
||||
@@ -292,7 +293,7 @@ class ValidationSelector(SOTAexpSelector):
|
||||
Sorts all valid experiments by score and returns the top N.
|
||||
"""
|
||||
|
||||
mock_folder = f"/tmp/mock/{self.competition}"
|
||||
mock_folder = f"/tmp/mock/{self.competition}" # nosec B108 — Docker volume mount point derived from internal competition name
|
||||
|
||||
try:
|
||||
data_py_code, grade_py_code = self._prepare_validation_scripts(
|
||||
@@ -539,7 +540,7 @@ def process_experiment(
|
||||
|
||||
# Run main script
|
||||
env = get_ds_env(
|
||||
extra_volumes={f"/tmp/mock/{competition}/{input_folder}": input_folder},
|
||||
extra_volumes={f"/tmp/mock/{competition}/{input_folder}": input_folder}, # nosec B108 — Docker volume mount point derived from internal competition name
|
||||
running_timeout_period=DS_RD_SETTING.full_timeout,
|
||||
)
|
||||
result = ws.run(env=env, entry="python main.py")
|
||||
@@ -587,8 +588,8 @@ def _parsing_score(grade_stdout: str) -> Optional[float]:
|
||||
except:
|
||||
pass
|
||||
try:
|
||||
# Priority 2: Eval dict
|
||||
return float(eval(json_str)["score"])
|
||||
# Priority 2: safe literal eval for Python-style dicts
|
||||
return float(ast.literal_eval(json_str)["score"])
|
||||
except:
|
||||
pass
|
||||
try:
|
||||
|
||||
@@ -35,10 +35,11 @@ def select(X: pd.DataFrame) -> pd.DataFrame:
|
||||
class KGModelFeatureSelectionCoder(Developer[KGModelExperiment]):
|
||||
def develop(self, exp: KGModelExperiment) -> KGModelExperiment:
|
||||
target_model_type = exp.sub_tasks[0].model_type
|
||||
assert target_model_type in KG_SELECT_MAPPING
|
||||
if target_model_type not in KG_SELECT_MAPPING:
|
||||
raise ValueError(f"target_model_type {target_model_type} not in KG_SELECT_MAPPING")
|
||||
if len(exp.experiment_workspace.data_description) == 1:
|
||||
code = (
|
||||
Environment(undefined=StrictUndefined)
|
||||
Environment(undefined=StrictUndefined) # nosec B701 — renders Python code templates, not HTML; autoescape would corrupt code
|
||||
.from_string(DEFAULT_SELECTION_CODE)
|
||||
.render(feature_index_list=None)
|
||||
)
|
||||
@@ -62,7 +63,7 @@ class KGModelFeatureSelectionCoder(Developer[KGModelExperiment]):
|
||||
chosen_index_to_list_index = [i - 1 for i in chosen_index]
|
||||
|
||||
code = (
|
||||
Environment(undefined=StrictUndefined)
|
||||
Environment(undefined=StrictUndefined) # nosec B701 — renders Python code templates, not HTML; autoescape would corrupt code
|
||||
.from_string(DEFAULT_SELECTION_CODE)
|
||||
.render(feature_index_list=chosen_index_to_list_index)
|
||||
)
|
||||
|
||||
@@ -165,7 +165,8 @@ class KGScenario(Scenario):
|
||||
return data_info
|
||||
|
||||
def output_format(self, tag=None) -> str:
|
||||
assert tag in [None, "feature", "model"]
|
||||
if tag not in [None, "feature", "model"]:
|
||||
raise ValueError(f"tag must be None, 'feature', or 'model', got {tag!r}")
|
||||
feature_output_format = f"""The feature code should output following the format:
|
||||
{T(".prompts:kg_feature_output_format").r()}"""
|
||||
model_output_format = f"""The model code should output following the format:\n""" + T(
|
||||
@@ -180,7 +181,8 @@ class KGScenario(Scenario):
|
||||
return model_output_format
|
||||
|
||||
def interface(self, tag=None) -> str:
|
||||
assert tag in [None, "feature", "XGBoost", "RandomForest", "LightGBM", "NN"]
|
||||
if tag not in [None, "feature", "XGBoost", "RandomForest", "LightGBM", "NN"]:
|
||||
raise ValueError(f"tag must be None, 'feature', 'XGBoost', 'RandomForest', 'LightGBM', or 'NN', got {tag!r}")
|
||||
feature_interface = f"""The feature code should follow the interface:
|
||||
{T(".prompts:kg_feature_interface").r()}"""
|
||||
if tag == "feature":
|
||||
@@ -195,7 +197,8 @@ class KGScenario(Scenario):
|
||||
return model_interface
|
||||
|
||||
def simulator(self, tag=None) -> str:
|
||||
assert tag in [None, "feature", "model"]
|
||||
if tag not in [None, "feature", "model"]:
|
||||
raise ValueError(f"tag must be None, 'feature', or 'model', got {tag!r}")
|
||||
|
||||
kg_feature_simulator = (
|
||||
"The feature code will be sent to the simulator:\n" + T(".prompts:kg_feature_simulator").r()
|
||||
|
||||
+6
-6
@@ -79,12 +79,12 @@ def preprocess_script():
|
||||
This method applies the preprocessing steps to the training, validation, and test datasets.
|
||||
"""
|
||||
if os.path.exists("/kaggle/input/X_train.pkl"):
|
||||
X_train = pd.read_pickle("/kaggle/input/X_train.pkl")
|
||||
X_valid = pd.read_pickle("/kaggle/input/X_valid.pkl")
|
||||
y_train = pd.read_pickle("/kaggle/input/y_train.pkl")
|
||||
y_valid = pd.read_pickle("/kaggle/input/y_valid.pkl")
|
||||
X_test = pd.read_pickle("/kaggle/input/X_test.pkl")
|
||||
others = pd.read_pickle("/kaggle/input/others.pkl")
|
||||
X_train = pd.read_pickle("/kaggle/input/X_train.pkl") # nosec B301 — trusted Kaggle input
|
||||
X_valid = pd.read_pickle("/kaggle/input/X_valid.pkl") # nosec B301
|
||||
y_train = pd.read_pickle("/kaggle/input/y_train.pkl") # nosec B301
|
||||
y_valid = pd.read_pickle("/kaggle/input/y_valid.pkl") # nosec B301
|
||||
X_test = pd.read_pickle("/kaggle/input/X_test.pkl") # nosec B301
|
||||
others = pd.read_pickle("/kaggle/input/others.pkl") # nosec B301
|
||||
y_train = pd.Series(y_train).reset_index(drop=True)
|
||||
y_valid = pd.Series(y_valid).reset_index(drop=True)
|
||||
|
||||
|
||||
+6
-6
@@ -85,12 +85,12 @@ def preprocess_script():
|
||||
This method applies the preprocessing steps to the training, validation, and test datasets.
|
||||
"""
|
||||
if os.path.exists("/kaggle/input/X_train.pkl"):
|
||||
X_train = pd.read_pickle("/kaggle/input/X_train.pkl")
|
||||
X_valid = pd.read_pickle("/kaggle/input/X_valid.pkl")
|
||||
y_train = pd.read_pickle("/kaggle/input/y_train.pkl")
|
||||
y_valid = pd.read_pickle("/kaggle/input/y_valid.pkl")
|
||||
X_test = pd.read_pickle("/kaggle/input/X_test.pkl")
|
||||
others = pd.read_pickle("/kaggle/input/others.pkl")
|
||||
X_train = pd.read_pickle("/kaggle/input/X_train.pkl") # nosec B301
|
||||
X_valid = pd.read_pickle("/kaggle/input/X_valid.pkl") # nosec B301
|
||||
y_train = pd.read_pickle("/kaggle/input/y_train.pkl") # nosec B301
|
||||
y_valid = pd.read_pickle("/kaggle/input/y_valid.pkl") # nosec B301
|
||||
X_test = pd.read_pickle("/kaggle/input/X_test.pkl") # nosec B301
|
||||
others = pd.read_pickle("/kaggle/input/others.pkl") # nosec B301
|
||||
|
||||
return X_train, X_valid, y_train, y_valid, X_test, *others
|
||||
X_train, X_valid, y_train, y_valid = prepreprocess()
|
||||
|
||||
+5
-5
@@ -82,11 +82,11 @@ def preprocess_script():
|
||||
This method applies the preprocessing steps to the training, validation, and test datasets.
|
||||
"""
|
||||
if os.path.exists("X_train.pkl"):
|
||||
X_train = pd.read_pickle("X_train.pkl")
|
||||
X_valid = pd.read_pickle("X_valid.pkl")
|
||||
y_train = pd.read_pickle("y_train.pkl")
|
||||
y_valid = pd.read_pickle("y_valid.pkl")
|
||||
X_test = pd.read_pickle("X_test.pkl")
|
||||
X_train = pd.read_pickle("X_train.pkl") # nosec B301
|
||||
X_valid = pd.read_pickle("X_valid.pkl") # nosec B301
|
||||
y_train = pd.read_pickle("y_train.pkl") # nosec B301
|
||||
y_valid = pd.read_pickle("y_valid.pkl") # nosec B301
|
||||
X_test = pd.read_pickle("X_test.pkl") # nosec B301
|
||||
return X_train, X_valid, y_train, y_valid, X_test
|
||||
|
||||
X_train, X_valid, y_train, y_valid, test, status_encoder, test_ids = prepreprocess()
|
||||
|
||||
+6
-6
@@ -73,12 +73,12 @@ def preprocess_script():
|
||||
This method applies the preprocessing steps to the training, validation, and test datasets.
|
||||
"""
|
||||
if os.path.exists("/kaggle/input/X_train.pkl"):
|
||||
X_train = pd.read_pickle("/kaggle/input/X_train.pkl")
|
||||
X_valid = pd.read_pickle("/kaggle/input/X_valid.pkl")
|
||||
y_train = pd.read_pickle("/kaggle/input/y_train.pkl")
|
||||
y_valid = pd.read_pickle("/kaggle/input/y_valid.pkl")
|
||||
X_test = pd.read_pickle("/kaggle/input/X_test.pkl")
|
||||
others = pd.read_pickle("/kaggle/input/others.pkl")
|
||||
X_train = pd.read_pickle("/kaggle/input/X_train.pkl") # nosec B301
|
||||
X_valid = pd.read_pickle("/kaggle/input/X_valid.pkl") # nosec B301
|
||||
y_train = pd.read_pickle("/kaggle/input/y_train.pkl") # nosec B301
|
||||
y_valid = pd.read_pickle("/kaggle/input/y_valid.pkl") # nosec B301
|
||||
X_test = pd.read_pickle("/kaggle/input/X_test.pkl") # nosec B301
|
||||
others = pd.read_pickle("/kaggle/input/others.pkl") # nosec B301
|
||||
y_train = pd.Series(y_train).reset_index(drop=True)
|
||||
y_valid = pd.Series(y_valid).reset_index(drop=True)
|
||||
|
||||
|
||||
@@ -87,7 +87,12 @@ def crawl_descriptions(
|
||||
content = e.get_attribute("innerHTML")
|
||||
contents.append(content)
|
||||
|
||||
assert len(subtitles) == len(contents) + 1 and subtitles[-1] == "Citation"
|
||||
if not (len(subtitles) == len(contents) + 1 and subtitles[-1] == "Citation"):
|
||||
raise AssertionError(
|
||||
f"Expected len(contents)+1 == len(subtitles) and last subtitle == 'Citation', "
|
||||
f"got len(subtitles)={len(subtitles)}, len(contents)={len(contents)}, "
|
||||
f"last subtitle={subtitles[-1]!r}"
|
||||
)
|
||||
for i in range(len(subtitles) - 1):
|
||||
descriptions[subtitles[i]] = contents[i]
|
||||
|
||||
|
||||
@@ -307,7 +307,8 @@ class KGHypothesisGen(FactorAndModelHypothesisGen):
|
||||
class KGHypothesis2Experiment(FactorAndModelHypothesis2Experiment):
|
||||
def prepare_context(self, hypothesis: Hypothesis, trace: Trace) -> Tuple[dict, bool]:
|
||||
scenario = trace.scen.get_scenario_all_desc(filtered_tag="hypothesis_and_experiment")
|
||||
assert isinstance(hypothesis, KGHypothesis)
|
||||
if not isinstance(hypothesis, KGHypothesis):
|
||||
raise TypeError("hypothesis must be an instance of KGHypothesis")
|
||||
experiment_output_format = (
|
||||
T("scenarios.kaggle.prompts:feature_experiment_output_format").r()
|
||||
if hypothesis.action in [KG_ACTION_FEATURE_ENGINEERING, KG_ACTION_FEATURE_PROCESSING]
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
import sys
|
||||
import os
|
||||
import logging
|
||||
from pathlib import Path
|
||||
"""
|
||||
Qlib Factor Runner - Executes factor backtests in Docker.
|
||||
@@ -29,6 +31,83 @@ from rdagent.scenarios.qlib.experiment.model_experiment import QlibModelExperime
|
||||
DIRNAME = Path(__file__).absolute().resolve().parent
|
||||
DIRNAME_local = Path.cwd()
|
||||
|
||||
|
||||
def _shift_daily_constant_factor_if_needed(factor_col: "pd.Series", factor_name: str) -> "pd.Series":
|
||||
"""Detect and fix look-ahead bias in daily-constant factors.
|
||||
|
||||
A factor is "daily-constant" when every minute bar within the same calendar
|
||||
day carries an identical value. This happens when LLM code computes a daily
|
||||
aggregate (e.g. today's log return) and forward-fills it across all intraday
|
||||
bars without shifting — meaning the end-of-day value is visible at 00:00.
|
||||
|
||||
Fix: shift by one trading day so that the value assigned to day T is the
|
||||
aggregate computed from day T-1, eliminating the forward-looking information.
|
||||
"""
|
||||
import numpy as np
|
||||
|
||||
try:
|
||||
notnull = factor_col.dropna()
|
||||
if len(notnull) < 200:
|
||||
return factor_col
|
||||
|
||||
datetimes = notnull.index.get_level_values("datetime")
|
||||
dates = datetimes.normalize()
|
||||
|
||||
# Sample up to 50 random days and check intra-day uniqueness
|
||||
unique_dates = pd.Series(dates.unique())
|
||||
sample_dates = unique_dates.sample(min(50, len(unique_dates)), random_state=42)
|
||||
|
||||
daily_unique_counts = []
|
||||
for d in sample_dates:
|
||||
mask = dates == d
|
||||
vals = notnull.values[mask]
|
||||
if len(vals) > 1:
|
||||
daily_unique_counts.append(len(np.unique(vals[~np.isnan(vals)])))
|
||||
|
||||
if not daily_unique_counts:
|
||||
return factor_col
|
||||
|
||||
# If >90% of sampled days have exactly 1 unique value → daily-constant
|
||||
fraction_constant = sum(1 for c in daily_unique_counts if c == 1) / len(daily_unique_counts)
|
||||
if fraction_constant < 0.90:
|
||||
return factor_col # Intraday factor — no shift needed
|
||||
|
||||
logger.warning(
|
||||
f"[LookAheadFix] Factor '{factor_name}' is daily-constant "
|
||||
f"({fraction_constant:.0%} of days). Applying 1-day shift to remove look-ahead bias."
|
||||
)
|
||||
|
||||
# Shift: for each instrument, map daily values forward by 1 trading day
|
||||
instruments = factor_col.index.get_level_values("instrument").unique()
|
||||
shifted_parts = []
|
||||
for inst in instruments:
|
||||
inst_series = factor_col.xs(inst, level="instrument")
|
||||
# Get one value per calendar day (the first non-null bar)
|
||||
inst_dt = inst_series.index.normalize()
|
||||
daily_vals = inst_series.groupby(inst_dt).first()
|
||||
# Shift by 1 day
|
||||
daily_vals_shifted = daily_vals.shift(1)
|
||||
# Forward-fill back to minute bars
|
||||
minute_idx = inst_series.index
|
||||
minute_dates = minute_idx.normalize()
|
||||
shifted_minute = minute_dates.map(daily_vals_shifted)
|
||||
shifted_s = pd.Series(
|
||||
shifted_minute.values,
|
||||
index=pd.MultiIndex.from_arrays(
|
||||
[inst_series.index, [inst] * len(inst_series)],
|
||||
names=["datetime", "instrument"],
|
||||
),
|
||||
name=factor_col.name,
|
||||
)
|
||||
shifted_parts.append(shifted_s)
|
||||
|
||||
return pd.concat(shifted_parts).sort_index()
|
||||
|
||||
except Exception as e:
|
||||
logger.debug(f"[LookAheadFix] Could not apply daily shift for '{factor_name}': {e}")
|
||||
return factor_col
|
||||
|
||||
|
||||
# TODO: supporting multiprocessing and keep previous results
|
||||
|
||||
|
||||
@@ -391,8 +470,19 @@ class QlibFactorRunner(CachedRunner[QlibFactorExperiment]):
|
||||
import numpy as np
|
||||
|
||||
try:
|
||||
# Get workspace path
|
||||
workspace_path = exp.experiment_workspace.workspace_path
|
||||
# Get workspace path — factor code and result.h5 live in sub_workspace_list[0],
|
||||
# not in experiment_workspace (which is the Qlib template workspace).
|
||||
workspace_path = None
|
||||
if exp.sub_workspace_list:
|
||||
for ws in exp.sub_workspace_list:
|
||||
if ws is not None and hasattr(ws, 'workspace_path'):
|
||||
candidate = ws.workspace_path / "result.h5"
|
||||
if candidate.exists():
|
||||
workspace_path = ws.workspace_path
|
||||
break
|
||||
if workspace_path is None:
|
||||
# Fallback to experiment_workspace
|
||||
workspace_path = exp.experiment_workspace.workspace_path
|
||||
if workspace_path is None:
|
||||
return None
|
||||
|
||||
@@ -409,6 +499,12 @@ class QlibFactorRunner(CachedRunner[QlibFactorExperiment]):
|
||||
factor_col = factor_values.iloc[:, 0]
|
||||
factor_name = factor_values.columns[0]
|
||||
|
||||
# Detect and fix look-ahead bias in daily-constant factors.
|
||||
# If a factor has the same value for all minute bars within each calendar day
|
||||
# it was computed from same-day data (e.g. today's close return at 00:00).
|
||||
# Fix: shift by 1 trading day so value at day T = aggregate of day T-1.
|
||||
factor_col = _shift_daily_constant_factor_if_needed(factor_col, factor_name)
|
||||
|
||||
# Load source data for forward returns
|
||||
data_path = (
|
||||
Path(__file__).parent.parent.parent.parent.parent
|
||||
@@ -587,10 +683,12 @@ class QlibFactorRunner(CachedRunner[QlibFactorExperiment]):
|
||||
from pathlib import Path
|
||||
from rdagent.components.backtesting import ResultsDatabase
|
||||
|
||||
# Get factor name from hypothesis
|
||||
# Get factor name: prefer hypothesis, fallback to result Series 'factor_name' key
|
||||
factor_name = "unknown"
|
||||
if hasattr(exp, 'hypothesis') and exp.hypothesis is not None:
|
||||
factor_name = getattr(exp.hypothesis, 'hypothesis', 'unknown')
|
||||
if factor_name == 'unknown' and isinstance(result, pd.Series) and 'factor_name' in result.index:
|
||||
factor_name = str(result['factor_name'])
|
||||
|
||||
# Check if already rejected by protection
|
||||
if getattr(exp, 'rejected_by_protection', False):
|
||||
@@ -824,41 +922,74 @@ class QlibFactorRunner(CachedRunner[QlibFactorExperiment]):
|
||||
"""
|
||||
Save factor time-series values as parquet for strategy building.
|
||||
|
||||
This is essential for walk-forward validation and strategy combination.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
factor_name : str
|
||||
Name of the factor
|
||||
exp : QlibFactorExperiment
|
||||
The experiment with factor values
|
||||
Reruns the factor code on the FULL 6-year dataset so the parquet covers
|
||||
the complete backtest range (not just the debug 2024 subset).
|
||||
"""
|
||||
import os as _os
|
||||
import subprocess
|
||||
import shutil
|
||||
import tempfile
|
||||
|
||||
try:
|
||||
# Get workspace path
|
||||
workspace_path = exp.experiment_workspace.workspace_path
|
||||
# factor.py lives in sub_workspace_list[0], not experiment_workspace
|
||||
workspace_path = None
|
||||
if exp.sub_workspace_list:
|
||||
for ws in exp.sub_workspace_list:
|
||||
if ws is not None and hasattr(ws, 'workspace_path'):
|
||||
fp = ws.workspace_path / "factor.py"
|
||||
if fp.exists():
|
||||
workspace_path = ws.workspace_path
|
||||
break
|
||||
if workspace_path is None:
|
||||
workspace_path = exp.experiment_workspace.workspace_path
|
||||
if workspace_path is None:
|
||||
return
|
||||
|
||||
result_h5 = workspace_path / "result.h5"
|
||||
if not result_h5.exists():
|
||||
factor_py = workspace_path / "factor.py"
|
||||
if not factor_py.exists():
|
||||
return
|
||||
|
||||
# Read factor values
|
||||
project_root = Path(__file__).parent.parent.parent.parent.parent
|
||||
full_data = (
|
||||
project_root
|
||||
/ "git_ignore_folder"
|
||||
/ "factor_implementation_source_data"
|
||||
/ "intraday_pv.h5"
|
||||
)
|
||||
if not full_data.exists():
|
||||
return
|
||||
|
||||
# Run factor code on full data in a temp workspace
|
||||
import pandas as pd
|
||||
df = pd.read_hdf(str(result_h5), key="data")
|
||||
with tempfile.TemporaryDirectory(prefix="predix_fullval_") as tmp_dir:
|
||||
tmp = Path(tmp_dir)
|
||||
shutil.copy(str(factor_py), str(tmp / "factor.py"))
|
||||
shutil.copy(str(full_data), str(tmp / "intraday_pv.h5"))
|
||||
|
||||
ret = subprocess.run(
|
||||
["sys.executable", "factor.py"],
|
||||
cwd=str(tmp),
|
||||
capture_output=True,
|
||||
timeout=300,
|
||||
)
|
||||
if ret.returncode != 0:
|
||||
# Fall back to debug-data result if full-data run fails
|
||||
result_h5 = workspace_path / "result.h5"
|
||||
if not result_h5.exists():
|
||||
return
|
||||
df = pd.read_hdf(str(result_h5), key="data")
|
||||
else:
|
||||
result_h5_full = tmp / "result.h5"
|
||||
if not result_h5_full.exists():
|
||||
return
|
||||
df = pd.read_hdf(str(result_h5_full), key="data")
|
||||
|
||||
if df is None or df.empty:
|
||||
return
|
||||
|
||||
# Get the factor series (first column)
|
||||
series = df.iloc[:, 0]
|
||||
series.name = factor_name
|
||||
|
||||
# Save to results/factors/values/
|
||||
project_root = Path(__file__).parent.parent.parent.parent.parent
|
||||
|
||||
# Parallel run isolation
|
||||
parallel_run_id = _os.getenv("PARALLEL_RUN_ID", "0")
|
||||
if parallel_run_id != "0":
|
||||
values_dir = project_root / "results" / "runs" / f"run{parallel_run_id}" / "factors" / "values"
|
||||
@@ -866,17 +997,12 @@ class QlibFactorRunner(CachedRunner[QlibFactorExperiment]):
|
||||
values_dir = project_root / "results" / "factors" / "values"
|
||||
|
||||
values_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Safe filename
|
||||
safe_name = factor_name.replace("/", "_").replace("\\", "_").replace(" ", "_")[:100]
|
||||
parquet_path = values_dir / f"{safe_name}.parquet"
|
||||
series.to_frame().to_parquet(str(parquet_path))
|
||||
|
||||
# Save as parquet (with datetime index)
|
||||
series.to_parquet(str(parquet_path))
|
||||
|
||||
except Exception as e:
|
||||
# Don't let factor value saving break the main workflow
|
||||
pass
|
||||
except Exception:
|
||||
logging.debug("Error in save_factor_values_to_parquet", exc_info=True)
|
||||
|
||||
def _log_result_warnings(self, factor_name: str, result, metrics: dict) -> None:
|
||||
"""
|
||||
|
||||
@@ -241,6 +241,7 @@ class StrategyBuilder:
|
||||
if data.get("status") == "success" and data.get("ic") is not None:
|
||||
factors.append(data)
|
||||
except Exception:
|
||||
logger.warning("Failed to load factor file %s", f, exc_info=True)
|
||||
continue
|
||||
|
||||
# Sort by absolute IC
|
||||
|
||||
@@ -30,7 +30,8 @@ def _build_execute_calls(exp: QlibFactorExperiment, base_feature_workspaces: lis
|
||||
execute_calls = []
|
||||
|
||||
if exp.sub_tasks:
|
||||
assert isinstance(exp.prop_dev_feedback, CoSTEERMultiFeedback)
|
||||
if not isinstance(exp.prop_dev_feedback, CoSTEERMultiFeedback):
|
||||
raise TypeError("exp.prop_dev_feedback must be of type CoSTEERMultiFeedback")
|
||||
execute_calls.extend(
|
||||
(implementation.execute, ("All",))
|
||||
for implementation, feedback in zip(exp.sub_workspace_list, exp.prop_dev_feedback)
|
||||
|
||||
@@ -23,14 +23,25 @@ $low: low price at 1-minute bar.
|
||||
$volume: volume at 1-minute bar (tick volume for FX).
|
||||
|
||||
## Important Notes for 1min Data
|
||||
- 96 bars = 1 trading day (24 hours for FX)
|
||||
- 1 bar = 1 minute (confirmed)
|
||||
- 16 bars = 16 minutes
|
||||
- 4 bars = 4 minutes
|
||||
- 1 bar = 1 minute
|
||||
- 60 bars = 1 hour
|
||||
- ~1440 bars = 1 full trading day (FX trades nearly 24h, Mon 00:00 - Fri 22:00 UTC approx.)
|
||||
- Typical bars per calendar day: ~1200-1440 (varies by weekday, holidays have fewer)
|
||||
- Do NOT assume 96 bars/day — the actual count depends on the date
|
||||
- Data range: 2020-01-01 to 2026-03-20
|
||||
- Instrument: EURUSD
|
||||
- Timezone: UTC
|
||||
|
||||
## IMPORTANT: Bars per Day Correction
|
||||
The dataset has approximately 1440 bars per full trading day (1 bar = 1 minute, ~24h of FX trading).
|
||||
Some older documentation incorrectly stated "96 bars = 1 day" — this is WRONG. Always use:
|
||||
- 60 bars = 1 hour
|
||||
- 480 bars = 8 hours (London session 08:00-16:00 UTC)
|
||||
- 180 bars = 3 hours (London/NY overlap 13:00-16:00 UTC)
|
||||
Use datetime hour filtering (e.g., `df[df.index.get_level_values('datetime').hour.between(8, 15)]`)
|
||||
to select session bars — do NOT use bar-count offsets to define sessions.
|
||||
|
||||
## Session Times (UTC)
|
||||
- Asian: 00:00-08:00 UTC (low volatility)
|
||||
- London: 08:00-16:00 UTC (high volatility)
|
||||
|
||||
@@ -104,7 +104,7 @@ qlib_factor_strategy: |-
|
||||
result_df.columns = ['daily_volume_price_divergence']
|
||||
```
|
||||
|
||||
4. **Process ALL data — do not filter dates**: The source HDF5 contains data from 2020-01-01 to 2026-03-20. Do NOT filter to a single year. If your output has only 314 entries (one year of daily data), the factor will be rejected. Expected output: ~1500+ daily entries for 2020-2026.
|
||||
4. **Process ALL data — do not filter dates**: The source HDF5 contains data from 2020-01-01 to 2026-03-20 (development runs may use a 2024-only debug dataset with ~300 entries, which is acceptable). Do NOT filter to a single year in your code. Write your code to process whatever date range is available in the HDF5 file — do not hardcode date filters. Expected output for production data: ~1500+ daily entries for 2020-2026. Expected output for debug data: ~300 daily entries for 2024. Both are valid.
|
||||
|
||||
5. **Use `transform()` instead of `apply()` for per-group calculations**: `transform()` preserves the original index while `apply()` may reduce the number of rows unexpectedly:
|
||||
```python
|
||||
@@ -121,6 +121,35 @@ qlib_factor_strategy: |-
|
||||
assert result_df.index.names == ['datetime', 'instrument'], f"Index names must be ['datetime', 'instrument'], got {result_df.index.names}"
|
||||
```
|
||||
|
||||
7. **NEVER use same-day aggregations as the factor value — always shift by 1 day**: If your factor computes a daily aggregate (e.g. daily close return, daily OHLC range, daily volume), that aggregate is only known at end-of-day. Using it at the start of the same day is look-ahead bias. You MUST shift the daily aggregate by 1 day before forward-filling to minute bars:
|
||||
```python
|
||||
# WRONG: look-ahead bias! Today's close return is not known at 00:00
|
||||
daily_ret = df['$close'].groupby(level='instrument').resample('1D', level='datetime').last().pct_change()
|
||||
result_df['my_factor'] = daily_ret.groupby(level='instrument').transform(lambda x: x.reindex(df.index.get_level_values('datetime'), method='ffill'))
|
||||
|
||||
# CORRECT: shift by 1 trading day so factor value at day T = aggregate of day T-1
|
||||
daily_close = df.groupby([df.index.get_level_values('datetime').normalize(), df.index.get_level_values('instrument')])['$close'].last()
|
||||
daily_close.index.names = ['date', 'instrument']
|
||||
daily_ret = daily_close.groupby(level='instrument').pct_change().shift(1) # <-- shift(1) is MANDATORY
|
||||
# then map back to minute bars via ffill
|
||||
```
|
||||
This rule applies to ALL daily aggregations: returns, OHLC stats, volume, momentum, slopes, etc.
|
||||
**Session-based aggregations (London, NY, Asian session returns) are also daily aggregations** — the London
|
||||
session (08:00-16:00 UTC) ends at 16:00, so its return must be shifted by 1 day before use.
|
||||
Intraday rolling factors (e.g. 30-min rolling std computed at bar t using only bars t-N..t-1) do NOT need this shift.
|
||||
|
||||
8. **PREFER pure intraday rolling factors**: Factors that use only a trailing window of recent bars (e.g.
|
||||
rolling(30).mean() of returns, RSI(14), Bollinger Band z-score) have NO look-ahead risk and vary every
|
||||
minute. These are the best candidates for short-horizon (60-180 bar) prediction. Examples:
|
||||
- Rolling 15-min / 30-min / 60-min return momentum (15, 30, 60 bars respectively)
|
||||
- Rolling volatility (std of returns over 20-60 bars)
|
||||
- Distance of close from N-bar moving average (z-score)
|
||||
- RSI or similar oscillators computed on 1-min bars
|
||||
- VWAP deviation (requires volume — use $volume column)
|
||||
Always use `.shift(1)` on the lagged window (e.g. `rolling(N).mean().shift(1)`) to avoid using the
|
||||
current bar's own price in its own feature value.
|
||||
NOTE: 1 bar = 1 minute. The data has ~1440 bars per full trading day. Do NOT use 96 as a day proxy.
|
||||
|
||||
qlib_factor_output_format: |-
|
||||
Your output should be a pandas dataframe similar to the following example information:
|
||||
<class 'pandas.core.frame.DataFrame'>
|
||||
|
||||
@@ -56,7 +56,8 @@ class QlibQuantScenario(Scenario):
|
||||
)
|
||||
|
||||
def background(self, tag=None) -> str:
|
||||
assert tag in [None, "factor", "model"]
|
||||
if tag not in [None, "factor", "model"]:
|
||||
raise ValueError(f"tag must be None, 'factor', or 'model', got {tag!r}")
|
||||
quant_background = "The background of the scenario is as follows:\n" + T(".prompts:qlib_quant_background").r(
|
||||
runtime_environment=self.get_runtime_environment(),
|
||||
)
|
||||
@@ -83,7 +84,8 @@ class QlibQuantScenario(Scenario):
|
||||
return self._source_data
|
||||
|
||||
def output_format(self, tag=None) -> str:
|
||||
assert tag in [None, "factor", "model"]
|
||||
if tag not in [None, "factor", "model"]:
|
||||
raise ValueError(f"tag must be None, 'factor', or 'model', got {tag!r}")
|
||||
factor_output_format = (
|
||||
"The factor code should output the following format:\n" + T(".prompts:qlib_factor_output_format").r()
|
||||
)
|
||||
@@ -99,7 +101,8 @@ class QlibQuantScenario(Scenario):
|
||||
return model_output_format
|
||||
|
||||
def interface(self, tag=None) -> str:
|
||||
assert tag in [None, "factor", "model"]
|
||||
if tag not in [None, "factor", "model"]:
|
||||
raise ValueError(f"tag must be None, 'factor', or 'model', got {tag!r}")
|
||||
factor_interface = (
|
||||
"The factor code should be written in the following interface:\n" + T(".prompts:qlib_factor_interface").r()
|
||||
)
|
||||
@@ -115,7 +118,8 @@ class QlibQuantScenario(Scenario):
|
||||
return model_interface
|
||||
|
||||
def simulator(self, tag=None) -> str:
|
||||
assert tag in [None, "factor", "model"]
|
||||
if tag not in [None, "factor", "model"]:
|
||||
raise ValueError(f"tag must be None, 'factor', or 'model', got {tag!r}")
|
||||
factor_simulator = "The factor code will be sent to the simulator:\n" + T(".prompts:qlib_factor_simulator").r()
|
||||
model_simulator = "The model code will be sent to the simulator:\n" + T(".prompts:qlib_model_simulator").r()
|
||||
|
||||
@@ -185,7 +189,8 @@ class QlibQuantScenario(Scenario):
|
||||
return common_description(action) + interface(action) + output(action) + simulator(action)
|
||||
|
||||
def get_runtime_environment(self, tag: str = None) -> str:
|
||||
assert tag in [None, "factor", "model"]
|
||||
if tag not in [None, "factor", "model"]:
|
||||
raise ValueError(f"tag must be None, 'factor', or 'model', got {tag!r}")
|
||||
|
||||
if tag is None or tag == "factor":
|
||||
# Use factor env to get the runtime environment
|
||||
|
||||
@@ -4,7 +4,7 @@ import shutil
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
from jinja2 import Environment, StrictUndefined
|
||||
from jinja2 import Environment, StrictUndefined, select_autoescape
|
||||
|
||||
from rdagent.components.coder.factor_coder.config import FACTOR_COSTEER_SETTINGS
|
||||
from rdagent.utils.env import QTDockerEnv
|
||||
@@ -21,14 +21,16 @@ def generate_data_folder_from_qlib():
|
||||
entry=f"python generate.py",
|
||||
)
|
||||
|
||||
assert (Path(__file__).parent / "factor_data_template" / "intraday_pv_all.h5").exists(), (
|
||||
"intraday_pv_all.h5 is not generated. It means rdagent/scenarios/qlib/experiment/factor_data_template/generate.py is not executed correctly. Please check the log: \n"
|
||||
+ execute_log
|
||||
)
|
||||
assert (Path(__file__).parent / "factor_data_template" / "intraday_pv_debug.h5").exists(), (
|
||||
"intraday_pv_debug.h5 is not generated. It means rdagent/scenarios/qlib/experiment/factor_data_template/generate.py is not executed correctly. Please check the log: \n"
|
||||
+ execute_log
|
||||
)
|
||||
if not (Path(__file__).parent / "factor_data_template" / "intraday_pv_all.h5").exists():
|
||||
raise FileNotFoundError(
|
||||
"intraday_pv_all.h5 is not generated. It means rdagent/scenarios/qlib/experiment/factor_data_template/generate.py is not executed correctly. Please check the log: \n"
|
||||
+ execute_log
|
||||
)
|
||||
if not (Path(__file__).parent / "factor_data_template" / "intraday_pv_debug.h5").exists():
|
||||
raise FileNotFoundError(
|
||||
"intraday_pv_debug.h5 is not generated. It means rdagent/scenarios/qlib/experiment/factor_data_template/generate.py is not executed correctly. Please check the log: \n"
|
||||
+ execute_log
|
||||
)
|
||||
|
||||
Path(FACTOR_COSTEER_SETTINGS.data_folder).mkdir(parents=True, exist_ok=True)
|
||||
shutil.copy(
|
||||
@@ -67,7 +69,7 @@ def get_file_desc(p: Path, variable_list=[]) -> str:
|
||||
"""
|
||||
p = Path(p)
|
||||
|
||||
JJ_TPL = Environment(undefined=StrictUndefined).from_string("""
|
||||
JJ_TPL = Environment(undefined=StrictUndefined, autoescape=select_autoescape()).from_string("""
|
||||
# {{file_name}}
|
||||
|
||||
## File Type
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
import logging
|
||||
import json
|
||||
import os
|
||||
from typing import List, Tuple
|
||||
|
||||
from rdagent.components.coder.factor_coder.factor import FactorExperiment, FactorTask
|
||||
@@ -9,6 +11,47 @@ from rdagent.scenarios.qlib.experiment.model_experiment import QlibModelExperime
|
||||
from rdagent.scenarios.qlib.experiment.quant_experiment import QlibQuantScenario
|
||||
from rdagent.utils.agent.tpl import T
|
||||
|
||||
|
||||
def _build_compressed_history(trace: Trace, max_history: int) -> str:
|
||||
"""Return hypothesis_and_feedback string with only `max_history` entries.
|
||||
|
||||
Older entries beyond the last 2 are compressed to one bullet line each.
|
||||
"""
|
||||
if len(trace.hist) == 0:
|
||||
return "No previous hypothesis and feedback available since it's the first round."
|
||||
|
||||
FULL_DETAIL = 2
|
||||
old_hist = trace.hist[:-FULL_DETAIL] if len(trace.hist) > FULL_DETAIL else []
|
||||
recent_hist = trace.hist[-FULL_DETAIL:] if len(trace.hist) > FULL_DETAIL else trace.hist
|
||||
|
||||
parts = []
|
||||
if old_hist:
|
||||
lines = ["## Earlier experiments (summarized):"]
|
||||
for exp, fb in old_hist:
|
||||
names = []
|
||||
for task in exp.sub_tasks:
|
||||
if task is not None and hasattr(task, "factor_name"):
|
||||
names.append(task.factor_name)
|
||||
elif task is not None and hasattr(task, "model_type"):
|
||||
names.append(getattr(task, "model_type", "model"))
|
||||
ic_str = ""
|
||||
try:
|
||||
if exp.result is not None and "IC" in exp.result.index:
|
||||
ic_str = f" IC={exp.result.loc['IC']:.4f}"
|
||||
except Exception:
|
||||
logging.debug("Exception caught", exc_info=True)
|
||||
decision = "PASS" if fb.decision else "FAIL"
|
||||
obs = (fb.observations or "")[:120].replace("\n", " ")
|
||||
lines.append(f"- [{decision}]{ic_str} {', '.join(names) or 'unknown'}: {obs}")
|
||||
parts.append("\n".join(lines))
|
||||
|
||||
if recent_hist:
|
||||
rt = Trace(trace.scen)
|
||||
rt.hist = recent_hist
|
||||
parts.append(T("scenarios.qlib.prompts:hypothesis_and_feedback").r(trace=rt))
|
||||
|
||||
return "\n\n".join(parts)
|
||||
|
||||
QlibFactorHypothesis = Hypothesis
|
||||
|
||||
|
||||
@@ -17,13 +60,10 @@ class QlibFactorHypothesisGen(FactorHypothesisGen):
|
||||
super().__init__(scen)
|
||||
|
||||
def prepare_context(self, trace: Trace) -> Tuple[dict, bool]:
|
||||
hypothesis_and_feedback = (
|
||||
T("scenarios.qlib.prompts:hypothesis_and_feedback").r(
|
||||
trace=trace,
|
||||
)
|
||||
if len(trace.hist) > 0
|
||||
else "No previous hypothesis and feedback available since it's the first round."
|
||||
)
|
||||
max_h = int(os.environ.get("QLIB_QUANT_MAX_FACTOR_HISTORY", "20"))
|
||||
limited = Trace(trace.scen)
|
||||
limited.hist = trace.hist[-max_h:] if len(trace.hist) > max_h else trace.hist
|
||||
hypothesis_and_feedback = _build_compressed_history(limited, max_h)
|
||||
last_hypothesis_and_feedback = (
|
||||
T("scenarios.qlib.prompts:last_hypothesis_and_feedback").r(
|
||||
experiment=trace.hist[-1][0], feedback=trace.hist[-1][1]
|
||||
@@ -70,15 +110,15 @@ class QlibFactorHypothesis2Experiment(FactorHypothesis2Experiment):
|
||||
if len(trace.hist) == 0:
|
||||
hypothesis_and_feedback = "No previous hypothesis and feedback available since it's the first round."
|
||||
else:
|
||||
max_h = int(os.environ.get("QLIB_QUANT_MAX_FACTOR_HISTORY", "20"))
|
||||
factor_hist = [
|
||||
e for e in trace.hist
|
||||
if not hasattr(e[0].hypothesis, "action") or e[0].hypothesis.action == "factor"
|
||||
][-max_h:]
|
||||
specific_trace = Trace(trace.scen)
|
||||
for i in range(len(trace.hist) - 1, -1, -1):
|
||||
if not hasattr(trace.hist[i][0].hypothesis, "action") or trace.hist[i][0].hypothesis.action == "factor":
|
||||
specific_trace.hist.insert(0, trace.hist[i])
|
||||
if len(specific_trace.hist) > 0:
|
||||
specific_trace.hist.reverse()
|
||||
hypothesis_and_feedback = T("scenarios.qlib.prompts:hypothesis_and_feedback").r(
|
||||
trace=specific_trace,
|
||||
)
|
||||
specific_trace.hist = factor_hist
|
||||
if specific_trace.hist:
|
||||
hypothesis_and_feedback = _build_compressed_history(specific_trace, max_h)
|
||||
else:
|
||||
hypothesis_and_feedback = "No previous hypothesis and feedback available."
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
@@ -152,9 +153,41 @@ class QlibQuantHypothesisGen(FactorAndModelHypothesisGen):
|
||||
factor_inserted = True
|
||||
if len(specific_trace.hist) > 0:
|
||||
specific_trace.hist.reverse()
|
||||
hypothesis_and_feedback = T("scenarios.qlib.prompts:hypothesis_and_feedback").r(
|
||||
trace=specific_trace,
|
||||
)
|
||||
# Keep only the 2 most recent experiments in full detail; compress older ones
|
||||
# to brief bullet points to stay within the LLM context window.
|
||||
FULL_DETAIL_COUNT = 2
|
||||
old_hist = specific_trace.hist[:-FULL_DETAIL_COUNT] if len(specific_trace.hist) > FULL_DETAIL_COUNT else []
|
||||
recent_hist = specific_trace.hist[-FULL_DETAIL_COUNT:] if len(specific_trace.hist) > FULL_DETAIL_COUNT else specific_trace.hist
|
||||
|
||||
parts = []
|
||||
if old_hist:
|
||||
summary_lines = ["## Earlier experiments (summarized):"]
|
||||
for exp, fb in old_hist:
|
||||
factor_names = []
|
||||
for task in exp.sub_tasks:
|
||||
if task is not None and hasattr(task, "factor_name"):
|
||||
factor_names.append(task.factor_name)
|
||||
elif task is not None and hasattr(task, "model_type"):
|
||||
factor_names.append(getattr(task, "model_type", "model"))
|
||||
names_str = ", ".join(factor_names) if factor_names else "unknown"
|
||||
ic_str = ""
|
||||
try:
|
||||
if exp.result is not None:
|
||||
ic_val = exp.result.loc["IC"] if "IC" in exp.result.index else ""
|
||||
ic_str = f" IC={ic_val:.4f}" if ic_val != "" else ""
|
||||
except Exception:
|
||||
logging.debug("Error getting IC", exc_info=True)
|
||||
decision_str = "PASS" if fb.decision else "FAIL"
|
||||
obs_short = (fb.observations or "")[:120].replace("\n", " ")
|
||||
summary_lines.append(f"- [{decision_str}]{ic_str} {names_str}: {obs_short}")
|
||||
parts.append("\n".join(summary_lines))
|
||||
|
||||
if recent_hist:
|
||||
recent_trace = Trace(specific_trace.scen)
|
||||
recent_trace.hist = recent_hist
|
||||
parts.append(T("scenarios.qlib.prompts:hypothesis_and_feedback").r(trace=recent_trace))
|
||||
|
||||
hypothesis_and_feedback = "\n\n".join(parts)
|
||||
else:
|
||||
hypothesis_and_feedback = "No previous hypothesis and feedback available."
|
||||
|
||||
|
||||
@@ -77,6 +77,7 @@ def count_valid_factors() -> int:
|
||||
if data.get("status") == "success" and data.get("ic") is not None:
|
||||
count += 1
|
||||
except Exception:
|
||||
logger.warning("Failed to load factor file %s", json_file, exc_info=True)
|
||||
continue
|
||||
|
||||
return count
|
||||
|
||||
@@ -71,7 +71,7 @@ def submit_for_grading(grading_url: str, model_path: str) -> dict | None:
|
||||
def main():
|
||||
MODEL_PATH = os.environ.get("MODEL_PATH")
|
||||
DATA_PATH = os.environ.get("DATA_PATH")
|
||||
OUTPUT_DIR = os.environ.get("OUTPUT_DIR", "/tmp/autorl_output")
|
||||
OUTPUT_DIR = os.environ.get("OUTPUT_DIR", "/tmp/autorl_output") # nosec B108 — Docker container output dir, configurable via env var
|
||||
GRADING_SERVER_URL = os.environ.get("GRADING_SERVER_URL", "")
|
||||
TRAIN_RATIO = float(os.environ.get("TRAIN_RATIO", "0.05"))
|
||||
NUM_EPOCHS = int(os.environ.get("NUM_EPOCHS", "3"))
|
||||
|
||||
@@ -391,7 +391,7 @@ def set_baseline():
|
||||
return jsonify({"baseline_score": score, "status": "set"})
|
||||
|
||||
|
||||
def run_server(task: str, base_model: str, workspace: str, host: str = "0.0.0.0", port: int = 5000):
|
||||
def run_server(task: str, base_model: str, workspace: str, host: str = "127.0.0.1", port: int = 5000):
|
||||
"""启动服务器"""
|
||||
init_server(task, base_model, workspace)
|
||||
logger.info(f"Grading Server | task={task} | {host}:{port}")
|
||||
@@ -435,7 +435,7 @@ class LocalServerContext(GradingServerContext):
|
||||
logger.info(f"[Local Mode] Starting evaluation server on port {self.port}...")
|
||||
self.server = init_server(self.task, self.base_model, self.workspace)
|
||||
|
||||
self._http_server = make_server("0.0.0.0", self.port, app, threaded=True)
|
||||
self._http_server = make_server("0.0.0.0", self.port, app, threaded=True) # nosec B104 — intentional: Docker sandbox requires all-interface binding
|
||||
self._thread = threading.Thread(target=self._http_server.serve_forever, daemon=True)
|
||||
self._thread.start()
|
||||
|
||||
@@ -488,7 +488,7 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--base-model", type=str, default="")
|
||||
parser.add_argument("--workspace", type=str, default=".")
|
||||
parser.add_argument("--port", type=int, default=5000)
|
||||
parser.add_argument("--host", type=str, default="0.0.0.0")
|
||||
parser.add_argument("--host", type=str, default="127.0.0.1")
|
||||
args = parser.parse_args()
|
||||
|
||||
run_server(args.task, args.base_model, args.workspace, args.host, args.port)
|
||||
|
||||
@@ -9,7 +9,7 @@ peft>=0.18.1
|
||||
|
||||
# Evaluation
|
||||
opencompass==0.5.1
|
||||
setuptools<75 # uv venv doesn't include, opencompass depends on pkg_resources
|
||||
setuptools>=78.1.1 # Security fix: GHSA-8g6x-3r52-4m6c (path traversal in PackageIndex.download, arbitrary file write/RCE)
|
||||
|
||||
# Inference acceleration (optional, TRL supports 0.10.2-0.12.0)
|
||||
# Security: Version >=0.14.0 fixes CVE-2026-22807 (RCE via auto_map dynamic module loading)
|
||||
|
||||
@@ -9,7 +9,7 @@ from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
from jinja2 import Environment, FunctionLoader, StrictUndefined
|
||||
from jinja2 import Environment, FunctionLoader, StrictUndefined, select_autoescape
|
||||
|
||||
from rdagent.core.conf import RD_AGENT_SETTINGS
|
||||
from rdagent.log import rdagent_logger as logger
|
||||
@@ -38,7 +38,8 @@ def load_content(uri: str, caller_dir: Path | None = None, ftype: str = "yaml")
|
||||
caller_dir = get_caller_dir(upshift=1)
|
||||
# Parse the URI
|
||||
path_part, *yaml_trace = uri.split(":")
|
||||
assert len(yaml_trace) <= 1, f"Invalid uri {uri}, only one yaml trace is allowed."
|
||||
if len(yaml_trace) > 1:
|
||||
raise ValueError(f"Invalid uri {uri}, only one yaml trace is allowed.")
|
||||
yaml_trace = [key for yt in yaml_trace for key in yt.split(".")]
|
||||
|
||||
# load file_path with priorities.
|
||||
@@ -126,7 +127,7 @@ class RDAT:
|
||||
# loader=FunctionLoader(load_conent) is for supporting grammar like below.
|
||||
# `{% include "scenarios.data_science.share:component_spec.DataLoadSpec" %}`
|
||||
rendered = (
|
||||
Environment(undefined=StrictUndefined, loader=FunctionLoader(load_content))
|
||||
Environment(undefined=StrictUndefined, loader=FunctionLoader(load_content), autoescape=select_autoescape())
|
||||
.from_string(self.template)
|
||||
.render(**context)
|
||||
.strip("\n")
|
||||
|
||||
+29
-26
@@ -614,13 +614,14 @@ class LocalEnv(Env[ASpecificLocalConf]):
|
||||
if self.conf.extra_volumes is not None:
|
||||
for lp, rp in self.conf.extra_volumes.items():
|
||||
volumes[lp] = rp["bind"] if isinstance(rp, dict) else rp
|
||||
cache_path = "/tmp/sample" if "/sample/" in "".join(self.conf.extra_volumes.keys()) else "/tmp/full"
|
||||
cache_path = "/tmp/sample" if "/sample/" in "".join(self.conf.extra_volumes.keys()) else "/tmp/full" # nosec B108 — fixed Docker volume mount point, not a user-writable temp file
|
||||
Path(cache_path).mkdir(parents=True, exist_ok=True)
|
||||
volumes[cache_path] = T("scenarios.data_science.share:scen.cache_path").r()
|
||||
for lp, rp in running_extra_volume.items():
|
||||
volumes[lp] = rp
|
||||
|
||||
assert local_path is not None, "local_path should not be None"
|
||||
if local_path is None:
|
||||
raise ValueError("local_path should not be None")
|
||||
volumes = normalize_volumes(volumes, local_path)
|
||||
|
||||
@contextlib.contextmanager
|
||||
@@ -678,7 +679,7 @@ class LocalEnv(Env[ASpecificLocalConf]):
|
||||
cwd = Path(local_path).resolve() if local_path else None
|
||||
env = {k: str(v) if isinstance(v, int) else v for k, v in env.items()}
|
||||
|
||||
process = subprocess.Popen(
|
||||
process = subprocess.Popen( # nosec B602 — entry is an internal command string set by LocalEnvConf, not user input
|
||||
entry,
|
||||
cwd=cwd,
|
||||
env={**os.environ, **env},
|
||||
@@ -761,12 +762,15 @@ class CondaConf(LocalConf):
|
||||
to ensure bin_path is set correctly even if the conda env was just created.
|
||||
"""
|
||||
conda_path_result = subprocess.run(
|
||||
f"conda run -n {self.conda_env_name} --no-capture-output env | grep '^PATH='",
|
||||
["conda", "run", "-n", self.conda_env_name, "--no-capture-output", "env"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
shell=True,
|
||||
)
|
||||
self.bin_path = conda_path_result.stdout.strip().split("=")[1] if conda_path_result.returncode == 0 else ""
|
||||
if conda_path_result.returncode == 0:
|
||||
path_lines = [l for l in conda_path_result.stdout.splitlines() if l.startswith("PATH=")]
|
||||
self.bin_path = path_lines[0].split("=", 1)[1] if path_lines else ""
|
||||
else:
|
||||
self.bin_path = ""
|
||||
|
||||
|
||||
class MLECondaConf(CondaConf):
|
||||
@@ -850,24 +854,22 @@ class QlibCondaEnv(LocalEnv[QlibCondaConf]):
|
||||
def prepare(self) -> None:
|
||||
"""Prepare the conda environment if not already created."""
|
||||
try:
|
||||
envs = subprocess.run("conda env list", capture_output=True, text=True, shell=True)
|
||||
envs = subprocess.run(["conda", "env", "list"], capture_output=True, text=True)
|
||||
if self.conf.conda_env_name not in envs.stdout:
|
||||
print(f"[yellow]Conda env '{self.conf.conda_env_name}' not found, creating...[/yellow]")
|
||||
subprocess.check_call(
|
||||
f"conda create -y -n {self.conf.conda_env_name} python=3.10",
|
||||
shell=True,
|
||||
["conda", "create", "-y", "-n", self.conf.conda_env_name, "python=3.10"],
|
||||
)
|
||||
subprocess.check_call(
|
||||
f"conda run -n {self.conf.conda_env_name} pip install --upgrade pip cython",
|
||||
shell=True,
|
||||
["conda", "run", "-n", self.conf.conda_env_name, "pip", "install", "--upgrade", "pip", "cython"],
|
||||
)
|
||||
subprocess.check_call(
|
||||
f"conda run -n {self.conf.conda_env_name} pip install git+https://github.com/microsoft/qlib.git@2fb9380b342556ddb50a4b24e4fe8655d548b2b8",
|
||||
shell=True,
|
||||
["conda", "run", "-n", self.conf.conda_env_name, "pip", "install",
|
||||
"git+https://github.com/microsoft/qlib.git@2fb9380b342556ddb50a4b24e4fe8655d548b2b8"],
|
||||
)
|
||||
subprocess.check_call(
|
||||
f"conda run -n {self.conf.conda_env_name} pip install catboost xgboost tables torch",
|
||||
shell=True,
|
||||
["conda", "run", "-n", self.conf.conda_env_name, "pip", "install",
|
||||
"catboost", "xgboost", "tables", "torch"],
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
@@ -888,10 +890,9 @@ def _sync_conda_cache_with_real_envs() -> None:
|
||||
"""Ensure the prepared cache includes environments that already exist on disk."""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
"conda env list",
|
||||
["conda", "env", "list"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
shell=True,
|
||||
check=False,
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - best-effort helper
|
||||
@@ -924,14 +925,15 @@ def _prepare_conda_env(env_name: str, requirements_file: Path, python_version: s
|
||||
python_version: Python version for the environment
|
||||
"""
|
||||
# 1. Create conda environment if not exists
|
||||
result = subprocess.run(f"conda env list | grep -q '^{env_name} '", shell=True)
|
||||
if result.returncode != 0:
|
||||
env_list = subprocess.run(["conda", "env", "list"], capture_output=True, text=True, check=False)
|
||||
env_exists = any(line.split()[0] == env_name for line in env_list.stdout.splitlines() if line and not line.startswith("#"))
|
||||
if not env_exists:
|
||||
print(f"[yellow]Creating conda env '{env_name}' (Python {python_version})...[/yellow]")
|
||||
subprocess.check_call(f"conda create -y -n {env_name} python={python_version}", shell=True)
|
||||
subprocess.check_call(f"conda run -n {env_name} pip install --upgrade pip", shell=True)
|
||||
subprocess.check_call(["conda", "create", "-y", "-n", env_name, f"python={python_version}"])
|
||||
subprocess.check_call(["conda", "run", "-n", env_name, "pip", "install", "--upgrade", "pip"])
|
||||
|
||||
print(f"[yellow]Installing dependencies from {requirements_file.name}...[/yellow]")
|
||||
subprocess.check_call(f"conda run -n {env_name} pip install -r {requirements_file}", shell=True)
|
||||
subprocess.check_call(["conda", "run", "-n", env_name, "pip", "install", "-r", str(requirements_file)])
|
||||
print(f"[green]Conda env '{env_name}' ready[/green]")
|
||||
|
||||
_CONDA_ENV_PREPARED.add(env_name)
|
||||
@@ -971,8 +973,8 @@ class FTCondaEnv(LocalEnv[FTCondaConf]):
|
||||
# Note: flash-attn>=2.8 is required for B200 (sm_100) support
|
||||
print("[yellow]Installing flash-attn (compiling, may take a few minutes)...[/yellow]")
|
||||
subprocess.check_call(
|
||||
f"conda run -n {self.conf.conda_env_name} pip install 'flash-attn>=2.8' --no-build-isolation --no-cache-dir",
|
||||
shell=True,
|
||||
["conda", "run", "-n", self.conf.conda_env_name, "pip", "install",
|
||||
"flash-attn>=2.8", "--no-build-isolation", "--no-cache-dir"],
|
||||
)
|
||||
|
||||
# Re-update bin_path after prepare() in case the conda env was just created
|
||||
@@ -1442,7 +1444,7 @@ class DockerEnv(Env[DockerConf]):
|
||||
if self.conf.extra_volumes is not None:
|
||||
for lp, rp in self.conf.extra_volumes.items():
|
||||
volumes[lp] = rp if isinstance(rp, dict) else {"bind": rp, "mode": self.conf.extra_volume_mode}
|
||||
cache_path = "/tmp/sample" if "/sample/" in "".join(self.conf.extra_volumes.keys()) else "/tmp/full"
|
||||
cache_path = "/tmp/sample" if "/sample/" in "".join(self.conf.extra_volumes.keys()) else "/tmp/full" # nosec B108 — fixed Docker volume mount point, not a user-writable temp file
|
||||
Path(cache_path).mkdir(parents=True, exist_ok=True)
|
||||
volumes[cache_path] = {
|
||||
"bind": T("scenarios.data_science.share:scen.cache_path").r(),
|
||||
@@ -1471,7 +1473,8 @@ class DockerEnv(Env[DockerConf]):
|
||||
cpu_count=self.conf.cpu_count, # Set CPU limit
|
||||
**self._gpu_kwargs(client),
|
||||
)
|
||||
assert container is not None # Ensure container was created successfully
|
||||
if container is None:
|
||||
raise AssertionError("Docker container was not created successfully")
|
||||
logs = container.logs(stream=True)
|
||||
print(Rule("[bold green]Docker Logs Begin[/bold green]", style="dark_orange"))
|
||||
table = Table(title="Run Info", show_header=False)
|
||||
|
||||
@@ -270,6 +270,11 @@ class LoopBase:
|
||||
msg = "We have reset the loop instance, stop all the routines and resume."
|
||||
raise self.LoopResumeError(msg) from e
|
||||
else:
|
||||
# Do NOT advance step_idx for unhandled exceptions (e.g. LoopResumeError
|
||||
# propagating from _propose). Keeping step_idx at the current step lets
|
||||
# kickoff_loop retry step 0 on the next resume instead of permanently
|
||||
# corrupting the loop with a missing direct_exp_gen result.
|
||||
step_forward = False
|
||||
raise # re-raise unhandled exceptions
|
||||
finally:
|
||||
# No matter the execution succeed or not, we have to finish the following steps
|
||||
|
||||
@@ -30,7 +30,8 @@ def wait_retry(
|
||||
>>> counter
|
||||
2
|
||||
"""
|
||||
assert retry_n > 0, "retry_n should be greater than 0"
|
||||
if retry_n <= 0:
|
||||
raise ValueError("retry_n should be greater than 0")
|
||||
|
||||
def decorator(f: Callable[..., ASpecificRet]) -> Callable[..., ASpecificRet]:
|
||||
def wrapper(*args: Any, **kwargs: Any) -> ASpecificRet:
|
||||
|
||||
@@ -84,7 +84,8 @@ class WorkflowTracker:
|
||||
# Log timer status if timer is started
|
||||
if self.loop_base.timer.started:
|
||||
remain_time = self.loop_base.timer.remain_time()
|
||||
assert remain_time is not None
|
||||
if remain_time is None:
|
||||
raise AssertionError("remain_time should not be None")
|
||||
mlflow.log_metric("remain_time", remain_time.total_seconds())
|
||||
mlflow.log_metric(
|
||||
"remain_percent",
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
{
|
||||
"release-type": "python",
|
||||
"bump-minor-pre-major": true,
|
||||
"bump-patch-for-minor-pre-major": true,
|
||||
"packages": {
|
||||
".": {
|
||||
"release-type": "python",
|
||||
"bump-minor-pre-major": true,
|
||||
"bump-patch-for-minor-pre-major": true
|
||||
}
|
||||
}
|
||||
}
|
||||
+5
-5
@@ -9,8 +9,8 @@ psutil
|
||||
fire
|
||||
fuzzywuzzy
|
||||
openai
|
||||
litellm>=1.73 # to support `from litellm import get_valid_models`
|
||||
aiohttp>=3.13.4 # CVE-2026-22815, CVE-2026-34515, CVE-2026-34516, CVE-2026-34525
|
||||
litellm>=1.83.14 # to support `from litellm import get_valid_models`
|
||||
aiohttp>=3.13.4 # CVE-2026-22815, CVE-2026-34515, CVE-2026-34516, CVE-2026-34525; >=3.13.4 due to litellm==1.83.14 exact pin
|
||||
azure.identity
|
||||
pyarrow
|
||||
rich
|
||||
@@ -37,7 +37,7 @@ tables
|
||||
tree-sitter-python
|
||||
tree-sitter
|
||||
|
||||
python-dotenv
|
||||
python-dotenv>=1.2.2 # CVE: symlink following allows arbitrary file overwrite
|
||||
|
||||
# infrastructure related.
|
||||
docker
|
||||
@@ -98,8 +98,8 @@ optuna>=3.5.0
|
||||
beautifulsoup4>=4.12.0
|
||||
|
||||
# ML Training Pipeline
|
||||
lightgbm>=3.3.0
|
||||
scipy>=1.9.0
|
||||
lightgbm>=3.3.5
|
||||
scipy>=1.15.3
|
||||
|
||||
# RL Trading (optional - system works without these)
|
||||
# Install for full RL training: pip install stable-baselines3[extra] gymnasium
|
||||
|
||||
+1
-1
@@ -8,7 +8,7 @@
|
||||
# Only install if you want to use full PPO/A2C/SAC training.
|
||||
|
||||
# Core RL library
|
||||
stable-baselines3[extra]>=2.0.0
|
||||
stable-baselines3[extra]>=2.8.0
|
||||
|
||||
# Gymnasium environment (OpenAI Gym successor)
|
||||
gymnasium>=0.29.0
|
||||
|
||||
@@ -198,6 +198,55 @@ def scan_factors(workspace_dir: Path, skip_evaluated: bool = True) -> List[Facto
|
||||
return factors
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Look-ahead bias detection for daily-constant factors
|
||||
# ---------------------------------------------------------------------------
|
||||
def _shift_daily_constant_factor_if_needed(factor_col: "pd.Series", factor_name: str) -> "pd.Series":
|
||||
"""Detect daily-constant factors (look-ahead bias) and shift by 1 trading day."""
|
||||
sample_days = factor_col.index.get_level_values("datetime").normalize().unique()
|
||||
if len(sample_days) < 10:
|
||||
return factor_col
|
||||
rng = np.random.default_rng(42)
|
||||
days_to_check = rng.choice(sample_days, size=min(50, len(sample_days)), replace=False)
|
||||
constant_count = 0
|
||||
for day in days_to_check:
|
||||
day_mask = factor_col.index.get_level_values("datetime").normalize() == day
|
||||
day_vals = factor_col[day_mask].dropna()
|
||||
if len(day_vals) == 0:
|
||||
continue
|
||||
if day_vals.nunique() == 1:
|
||||
constant_count += 1
|
||||
fraction_constant = constant_count / len(days_to_check)
|
||||
if fraction_constant < 0.90:
|
||||
return factor_col
|
||||
# Shift by 1 trading day per instrument
|
||||
import logging
|
||||
logging.getLogger(__name__).info(
|
||||
"Factor '%s' is %.0f%% daily-constant — shifting 1 trading day to fix look-ahead bias",
|
||||
factor_name, fraction_constant * 100,
|
||||
)
|
||||
instruments = factor_col.index.get_level_values("instrument").unique() if "instrument" in factor_col.index.names else [None]
|
||||
shifted_parts = []
|
||||
for instr in instruments:
|
||||
if instr is not None:
|
||||
mask = factor_col.index.get_level_values("instrument") == instr
|
||||
col_instr = factor_col[mask]
|
||||
else:
|
||||
col_instr = factor_col
|
||||
dates = col_instr.index.get_level_values("datetime").normalize()
|
||||
trading_days = dates.unique().sort_values()
|
||||
day_first = col_instr.groupby(dates).first()
|
||||
day_first_shifted = day_first.shift(1)
|
||||
day_first_shifted.index = pd.to_datetime(day_first_shifted.index)
|
||||
day_map = day_first_shifted.reindex(pd.to_datetime(trading_days)).values
|
||||
new_vals = pd.Series(
|
||||
day_map[np.searchsorted(trading_days.values, dates.values)],
|
||||
index=col_instr.index,
|
||||
)
|
||||
shifted_parts.append(new_vals)
|
||||
return pd.concat(shifted_parts).sort_index()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Factor evaluator
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -263,6 +312,7 @@ def evaluate_factor_full(factor: FactorInfo, full_data: pd.DataFrame,
|
||||
result = pd.read_hdf(str(result_file), key="data")
|
||||
total_count = len(result)
|
||||
factor_val = result.iloc[:, 0]
|
||||
factor_val = _shift_daily_constant_factor_if_needed(factor_val, factor.factor_name)
|
||||
non_null_count = factor_val.notna().sum()
|
||||
|
||||
if non_null_count < 1000:
|
||||
|
||||
@@ -55,7 +55,7 @@ if TRADING_STYLE == 'daytrading':
|
||||
FORWARD_BARS = int(os.getenv('FORWARD_BARS', '12'))
|
||||
MIN_IC = 0.02
|
||||
MIN_SHARPE = 0.5
|
||||
MIN_TRADES = 20
|
||||
MIN_TRADES = 300
|
||||
MAX_DRAWDOWN = -0.10
|
||||
STYLE_EMOJI = '🎯 Daytrading'
|
||||
STYLE_DESC = 'short-term intraday with FTMO compliance'
|
||||
@@ -68,9 +68,40 @@ else:
|
||||
STYLE_EMOJI = '📈 Swing'
|
||||
STYLE_DESC = 'medium-term intraday'
|
||||
|
||||
TXN_COST_BPS = float(os.getenv('TXN_COST_BPS', '1.0'))
|
||||
# Whether to use raw OHLCV-only strategies (no daily factors)
|
||||
OHLCV_ONLY = os.getenv('OHLCV_ONLY', '0') == '1'
|
||||
|
||||
console = Console()
|
||||
TXN_COST_BPS = float(os.getenv('TXN_COST_BPS', '2.14')) # 2.35 pip realistic EUR/USD costs
|
||||
|
||||
# ── Logging setup: everything printed goes to log file + stdout ───────────────
|
||||
_LOG_DIR = Path(__file__).parent.parent / "git_ignore_folder" / "logs"
|
||||
_LOG_DIR.mkdir(parents=True, exist_ok=True)
|
||||
_log_file_path = _LOG_DIR / f"gen_strategies_{datetime.now().strftime('%Y%m%d_%H%M%S')}.log"
|
||||
_log_file = open(_log_file_path, "w", encoding="utf-8", buffering=1) # line-buffered
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||||
handlers=[
|
||||
logging.StreamHandler(sys.stdout),
|
||||
logging.FileHandler(_log_file_path, encoding="utf-8"),
|
||||
],
|
||||
)
|
||||
|
||||
class _TeeFile:
|
||||
"""Writes to both stdout and log file — used as Rich Console file."""
|
||||
def __init__(self, *files):
|
||||
self._files = files
|
||||
def write(self, data):
|
||||
for f in self._files:
|
||||
f.write(data)
|
||||
def flush(self):
|
||||
for f in self._files:
|
||||
f.flush()
|
||||
def fileno(self):
|
||||
return self._files[0].fileno()
|
||||
|
||||
console = Console(file=_TeeFile(sys.stdout, _log_file), highlight=False)
|
||||
|
||||
# ============================================================================
|
||||
# LLM Configuration (Process-safe)
|
||||
@@ -153,7 +184,44 @@ def generate_single_strategy(args):
|
||||
factor_list = "\n".join([f"- {f['name']} (IC={f['ic']:.4f})" for f in factor_subset])
|
||||
|
||||
# Optimized prompts for daytrading vs swing
|
||||
if TRADING_STYLE == 'daytrading':
|
||||
if TRADING_STYLE == 'daytrading' and OHLCV_ONLY:
|
||||
system_prompt = """You are an expert EUR/USD intraday quant. You build strategies that work ONLY on raw price data (OHLCV), computing all indicators directly from the 1-minute close series.
|
||||
|
||||
CRITICAL RULES:
|
||||
1. The code receives ONLY a pandas Series called 'close' (1-minute EUR/USD close prices, UTC timestamps).
|
||||
2. 'factors' is NOT available — compute everything from 'close' directly.
|
||||
3. Create a pandas Series called 'signal' with values: 1 (long), -1 (short), 0 (neutral).
|
||||
4. signal.index MUST match close.index exactly.
|
||||
5. signal.name must be 'signal'.
|
||||
6. Use ONLY pandas/numpy — no external libraries.
|
||||
7. MANDATORY: The signal MUST flip at least 300 times across the full dataset. Use low thresholds.
|
||||
|
||||
Allowed intraday techniques (pick 2-3 and combine):
|
||||
- Session timing: London open (07:00-09:00 UTC), NY open (13:00-15:00 UTC), session overlap
|
||||
- Short-window RSI (7-14 bars) on 1-min close
|
||||
- EMA crossovers (fast=5-15 bars, slow=20-60 bars)
|
||||
- Bollinger Bands (20-bar, 1.5σ) for mean reversion
|
||||
- ATR-based volatility breakouts
|
||||
- VWAP deviation (approximate with rolling mean)
|
||||
- Time-of-day filters combined with momentum
|
||||
|
||||
Output ONLY valid JSON:
|
||||
{"strategy_name": "short_name", "factor_names": [], "description": "one sentence", "code": "python code"}"""
|
||||
|
||||
user_prompt = f"""Create a EUR/USD 1-minute intraday strategy using ONLY the raw close price series.
|
||||
|
||||
{f'Previous feedback: {feedback}' if feedback else 'First attempt — be creative and combine session timing with a momentum or mean-reversion indicator!'}
|
||||
|
||||
Hard requirements:
|
||||
- Signal must change direction at least 300 times total (~4-8 trades per trading day)
|
||||
- NEVER use ffill() or forward-fill on the signal — recompute fresh at every bar
|
||||
- Use RSI thresholds between 35-45 (long) and 55-65 (short) — NOT extreme values like 10/90
|
||||
- Use EMA crossover thresholds of 0 (cross above/below) for maximum trade frequency
|
||||
- Use causal indicators only: rolling windows, shift(1) — NO look-ahead bias
|
||||
- No factor data — compute everything from 'close'
|
||||
- Keep it simple: 2-3 indicators max"""
|
||||
|
||||
elif TRADING_STYLE == 'daytrading':
|
||||
system_prompt = f"""You are an expert daytrading quant specializing in EUR/USD scalping and intraday strategies.
|
||||
|
||||
CRITICAL RULES for {STYLE_DESC} (forward horizon: {FORWARD_BARS} bars = ~{FORWARD_BARS} minutes):
|
||||
@@ -162,27 +230,27 @@ CRITICAL RULES for {STYLE_DESC} (forward horizon: {FORWARD_BARS} bars = ~{FORWAR
|
||||
3. Create a pandas Series called 'signal' with values: 1 (long), -1 (short), 0 (neutral)
|
||||
4. signal.index MUST match close.index
|
||||
5. signal.name must be 'signal'
|
||||
6. Optimize for FREQUENT signals (many trades) since the horizon is only {FORWARD_BARS} minutes
|
||||
7. Use LOWER thresholds (0.2-0.5) to generate more trades for daytrading
|
||||
6. MANDATORY: signal must flip direction at least 300 times total — use low thresholds (0.1-0.3)
|
||||
7. Use rolling z-scores with SHORT windows (5-20 bars) and TIGHT thresholds
|
||||
|
||||
Output ONLY valid JSON with these fields:
|
||||
{{"strategy_name": "short_name", "factor_names": ["f1", "f2"], "description": "one sentence", "code": "python code"}}"""
|
||||
|
||||
|
||||
user_prompt = f"""Create a EUR/USD DAYTRADING strategy ({FORWARD_BARS}-minute horizon) using these factors:
|
||||
|
||||
{factor_list}
|
||||
|
||||
{f'Previous feedback: {feedback}' if feedback else 'First attempt - be creative!'}
|
||||
|
||||
Requirements for daytrading:
|
||||
- Use {FORWARD_BARS}-minute forward returns (not daily)
|
||||
- Generate frequent signals (aim for 20+ trades in the dataset)
|
||||
- Use rolling z-scores with short windows (10-30 bars)
|
||||
- Apply tight thresholds (0.2-0.5) for more trades
|
||||
- Combine momentum + mean-reversion effectively"""
|
||||
|
||||
Hard requirements:
|
||||
- signal must change at least 300 times total (~4 trades/day) — use thresholds of 0.1-0.3
|
||||
- NEVER use ffill() or forward-fill on the signal — recompute fresh at every bar
|
||||
- Use rolling z-scores with windows of 5-20 bars (not 50-100), thresholds ±0.2 to ±0.5
|
||||
- Combine 2 factors: one momentum, one mean-reversion
|
||||
- NO global mean/std — always use rolling(window).mean() with shift(1) to avoid look-ahead bias"""
|
||||
|
||||
else:
|
||||
system_prompt = f"""You are a quantitative trading expert specializing in EUR/USD intraday strategies.
|
||||
system_prompt = f"""You are a quantitative trading expert specializing in EUR/USD daily swing strategies.
|
||||
|
||||
CRITICAL RULES for {STYLE_DESC} (forward horizon: {FORWARD_BARS} bars = ~{FORWARD_BARS/60:.1f} hours):
|
||||
1. ONLY use the factors listed below - no others!
|
||||
@@ -190,15 +258,27 @@ CRITICAL RULES for {STYLE_DESC} (forward horizon: {FORWARD_BARS} bars = ~{FORWAR
|
||||
3. Create a pandas Series called 'signal' with values: 1 (long), -1 (short), 0 (neutral)
|
||||
4. signal.index MUST match close.index
|
||||
5. signal.name must be 'signal'
|
||||
6. IMPORTANT: factors are DAILY values broadcast to every 1-minute bar — they change once per day.
|
||||
Use daily-level logic: compare today's factor value to a rolling daily mean (window 5-20 DAYS).
|
||||
To get daily rolling mean: group by date, take first value per day, compute rolling, then reindex back.
|
||||
Example: dates = factors[col].index.get_level_values('datetime').normalize()
|
||||
daily_vals = factors[col].groupby(dates).first()
|
||||
daily_mean = daily_vals.rolling(10).mean().shift(1)
|
||||
daily_signal = (daily_vals > daily_mean).astype(int) * 2 - 1
|
||||
signal = daily_signal.reindex(dates).values (broadcast back to minute bars)
|
||||
7. The signal should change roughly once per day — this produces ~250-500 trades over 6 years.
|
||||
8. Keep conditions SIMPLE: one factor above/below its N-day rolling average. Avoid combining 3+ conditions.
|
||||
|
||||
Output ONLY valid JSON with these fields:
|
||||
{{"strategy_name": "short_name", "factor_names": ["f1", "f2"], "description": "one sentence", "code": "python code"}}"""
|
||||
|
||||
user_prompt = f"""Create a EUR/USD trading strategy using these factors:
|
||||
|
||||
user_prompt = f"""Create a EUR/USD SWING trading strategy (hold ~{FORWARD_BARS/60:.0f} hours) using these factors:
|
||||
|
||||
{factor_list}
|
||||
|
||||
{f'Previous feedback: {feedback}' if feedback else 'First attempt - be creative!'}"""
|
||||
{f'Previous feedback: {feedback}' if feedback else 'First attempt - be creative!'}
|
||||
|
||||
Use daily-level signal logic (factor above/below rolling daily mean). Signal changes once per day."""
|
||||
|
||||
api = APIBackend()
|
||||
response = api.build_messages_and_create_chat_completion(
|
||||
@@ -228,19 +308,27 @@ def run_backtest(close, factors_df, strategy_code):
|
||||
the signal, then delegate all metric computation to the unified
|
||||
``backtest_signal`` engine in the main process.
|
||||
"""
|
||||
if close is None or factors_df is None or len(factors_df.columns) < 2:
|
||||
if close is None:
|
||||
return None
|
||||
if not OHLCV_ONLY and (factors_df is None or len(factors_df.columns) < 2):
|
||||
return None
|
||||
|
||||
# Flatten MultiIndex — strategy code expects a plain DatetimeIndex
|
||||
if isinstance(close.index, pd.MultiIndex):
|
||||
close = close.droplevel(-1)
|
||||
close = close.sort_index()
|
||||
|
||||
import tempfile
|
||||
|
||||
# Subprocess stays minimal: it only runs the untrusted strategy code
|
||||
# and pickles the resulting signal. All numbers come from the shared engine.
|
||||
factors_line = "" if OHLCV_ONLY else "factors = pd.read_pickle('factors.pkl')"
|
||||
script = f"""
|
||||
import pandas as pd
|
||||
import numpy as np
|
||||
|
||||
close = pd.read_pickle('close.pkl')
|
||||
factors = pd.read_pickle('factors.pkl')
|
||||
{factors_line}
|
||||
|
||||
try:
|
||||
{chr(10).join(' ' + l for l in strategy_code.split(chr(10)))}
|
||||
@@ -258,7 +346,8 @@ signal.fillna(0).to_pickle('signal.pkl')
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
tdp = Path(td)
|
||||
close.to_pickle(str(tdp / 'close.pkl'))
|
||||
factors_df.to_pickle(str(tdp / 'factors.pkl'))
|
||||
if not OHLCV_ONLY and factors_df is not None:
|
||||
factors_df.to_pickle(str(tdp / 'factors.pkl'))
|
||||
(tdp / 'run.py').write_text(script)
|
||||
|
||||
try:
|
||||
@@ -276,27 +365,80 @@ signal.fillna(0).to_pickle('signal.pkl')
|
||||
except Exception as e:
|
||||
return {'status': 'failed', 'reason': str(e)[:200]}
|
||||
|
||||
# Main process: unified backtest (identical formulas everywhere).
|
||||
from rdagent.components.backtesting.vbt_backtest import backtest_signal
|
||||
# Main process: FTMO-realistic backtest (leverage + daily/total loss limits).
|
||||
from rdagent.components.backtesting.vbt_backtest import backtest_signal_ftmo
|
||||
|
||||
common = close.index.intersection(signal.index)
|
||||
if len(common) < 100:
|
||||
return {'status': 'failed', 'reason': f'Not enough aligned data ({len(common)} bars)'}
|
||||
|
||||
close_a = close.loc[common]
|
||||
close_a = close.loc[common]
|
||||
signal_a = signal.reindex(common).fillna(0)
|
||||
|
||||
# Forward returns at the configured horizon feed IC computation.
|
||||
fwd_returns = close_a.pct_change(FORWARD_BARS).shift(-FORWARD_BARS)
|
||||
|
||||
return backtest_signal(
|
||||
from rdagent.components.backtesting.vbt_backtest import OOS_START_DEFAULT
|
||||
return backtest_signal_ftmo(
|
||||
close=close_a,
|
||||
signal=signal_a,
|
||||
txn_cost_bps=TXN_COST_BPS,
|
||||
freq='1min',
|
||||
forward_returns=fwd_returns,
|
||||
oos_start=OOS_START_DEFAULT,
|
||||
wf_rolling=False, # too slow on 2M bars — run via rebacktest script instead
|
||||
mc_n_permutations=50,
|
||||
)
|
||||
|
||||
# ============================================================================
|
||||
# Threshold Tuner — relax numeric thresholds until MIN_TRADES is reached
|
||||
# ============================================================================
|
||||
def _rescale_thresholds(code: str, scale: float) -> str:
|
||||
"""
|
||||
Scale numeric literals in the strategy code that look like signal thresholds.
|
||||
RSI thresholds (30-70 range) are moved toward 50.
|
||||
Z-score / ratio thresholds (0.0–3.0 range) are multiplied by scale.
|
||||
"""
|
||||
import re
|
||||
|
||||
def replace_rsi(m):
|
||||
val = float(m.group(0))
|
||||
# Pull toward 50 by (1-scale) fraction
|
||||
new_val = 50 + (val - 50) * scale
|
||||
return f"{new_val:.1f}"
|
||||
|
||||
def replace_small(m):
|
||||
val = float(m.group(0))
|
||||
return f"{val * scale:.3f}"
|
||||
|
||||
# RSI-style thresholds: integers/floats between 10 and 90
|
||||
code = re.sub(r'\b([1-9]\d(?:\.\d+)?)\b', replace_rsi, code)
|
||||
# Small float thresholds: 0.05 – 2.99
|
||||
code = re.sub(r'\b(0\.\d+|[12]\.\d+)\b', replace_small, code)
|
||||
return code
|
||||
|
||||
|
||||
def tune_thresholds(close, factors_df, code: str) -> tuple:
|
||||
"""
|
||||
Binary-search scale factor (1.0 → 0.05) until n_trades >= MIN_TRADES.
|
||||
Returns (best_bt_result, tuned_code) where best_bt_result has max Sharpe
|
||||
among all runs that hit MIN_TRADES.
|
||||
"""
|
||||
best_bt, best_code = None, code
|
||||
|
||||
for scale in [1.0, 0.7, 0.5, 0.35, 0.2, 0.1, 0.05]:
|
||||
tuned = _rescale_thresholds(code, scale) if scale < 1.0 else code
|
||||
bt = run_backtest(close, factors_df, tuned)
|
||||
if bt is None or bt.get('status') != 'success':
|
||||
continue
|
||||
trades = bt.get('n_trades', 0)
|
||||
sharpe = bt.get('sharpe', -999)
|
||||
if trades >= MIN_TRADES:
|
||||
if best_bt is None or sharpe > best_bt.get('sharpe', -999):
|
||||
best_bt = bt
|
||||
best_code = tuned
|
||||
break # first scale that hits MIN_TRADES wins (they get looser after this)
|
||||
|
||||
return best_bt, best_code
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Main Parallel Strategy Generation
|
||||
# ============================================================================
|
||||
@@ -314,6 +456,7 @@ def main(target_count=10):
|
||||
)
|
||||
|
||||
console.print(f"\n[bold cyan]{STYLE_EMOJI} Parallel Strategy Generation[/bold cyan]")
|
||||
console.print(f"[dim]Log: {_log_file_path}[/dim]")
|
||||
console.print(f" Style: {STYLE_DESC}")
|
||||
console.print(f" Forward bars: {FORWARD_BARS}")
|
||||
console.print(f" Target: {target_count} accepted strategies")
|
||||
@@ -373,38 +516,83 @@ def main(target_count=10):
|
||||
if len(accepted) >= target_count:
|
||||
break
|
||||
|
||||
# Select random factor subset (2-5 factors)
|
||||
n_factors = random.randint(2, min(5, len(factors)))
|
||||
factor_subset = random.sample(factors, n_factors)
|
||||
|
||||
# Select random factor subset (2-5 factors) — empty for OHLCV-only mode
|
||||
if OHLCV_ONLY:
|
||||
factor_subset = []
|
||||
else:
|
||||
n_factors = random.randint(2, min(5, len(factors)))
|
||||
factor_subset = random.sample(factors, n_factors)
|
||||
|
||||
feedback = feedback_history[-1] if feedback_history and random.random() < 0.7 else None
|
||||
|
||||
|
||||
# Generate in main process (LLM doesn't parallelize well)
|
||||
gen_result = generate_single_strategy((attempt, factor_subset, feedback, attempt))
|
||||
|
||||
|
||||
if gen_result['status'] != 'generated':
|
||||
progress.update(task, advance=1)
|
||||
continue
|
||||
|
||||
|
||||
strategy = gen_result['strategy']
|
||||
|
||||
|
||||
# Backtest (main process - needs data access)
|
||||
# Build factors DataFrame for this strategy
|
||||
strat_factors = df_aligned[[f for f in strategy.get('factor_names', []) if f in df_aligned.columns]]
|
||||
if len(strat_factors.columns) < 2:
|
||||
progress.update(task, advance=1)
|
||||
continue
|
||||
|
||||
bt_result = run_backtest(close_aligned, strat_factors, strategy.get('code', ''))
|
||||
if OHLCV_ONLY:
|
||||
strat_factors = None
|
||||
bt_result = run_backtest(close, None, strategy.get('code', ''))
|
||||
else:
|
||||
strat_factors = df_aligned[[f for f in strategy.get('factor_names', []) if f in df_aligned.columns]]
|
||||
if len(strat_factors.columns) < 2:
|
||||
progress.update(task, advance=1)
|
||||
continue
|
||||
bt_result = run_backtest(close_aligned, strat_factors, strategy.get('code', ''))
|
||||
|
||||
if bt_result and bt_result.get('status') == 'success':
|
||||
ic = bt_result.get('ic', 0)
|
||||
sharpe = bt_result.get('sharpe', 0)
|
||||
trades = bt_result.get('n_trades', 0)
|
||||
dd = bt_result.get('max_drawdown', 0)
|
||||
|
||||
# Check acceptance criteria
|
||||
if abs(ic) > MIN_IC and sharpe > MIN_SHARPE and trades > MIN_TRADES and dd > MAX_DRAWDOWN:
|
||||
|
||||
# If too few trades, auto-tune thresholds before giving up
|
||||
original_code = strategy.get('code', '')
|
||||
if trades < MIN_TRADES and bt_result.get('status') == 'success':
|
||||
_log.info(f"TUNING trades={trades}<{MIN_TRADES} — trying looser thresholds")
|
||||
tuned_bt, tuned_code = tune_thresholds(
|
||||
close if OHLCV_ONLY else close_aligned,
|
||||
None if OHLCV_ONLY else strat_factors,
|
||||
original_code,
|
||||
)
|
||||
if tuned_bt and tuned_bt.get('n_trades', 0) >= MIN_TRADES:
|
||||
bt_result = tuned_bt
|
||||
strategy['code'] = tuned_code
|
||||
ic = bt_result.get('ic', 0)
|
||||
sharpe = bt_result.get('sharpe', 0)
|
||||
trades = bt_result.get('n_trades', 0)
|
||||
dd = bt_result.get('max_drawdown', 0)
|
||||
_log.info(f"TUNED Sharpe={sharpe:.2f} Trades={trades}")
|
||||
|
||||
# OOS metrics — mandatory, no fallback to IS values
|
||||
oos_sharpe = bt_result.get('oos_sharpe')
|
||||
oos_monthly = bt_result.get('oos_monthly_return_pct')
|
||||
oos_trades = bt_result.get('oos_n_trades', 0)
|
||||
|
||||
# Reject if OOS data is missing (strategy trained on data without OOS period)
|
||||
if oos_sharpe is None or oos_monthly is None:
|
||||
_log.info(f"REJECTED no OOS data (data ends before {OOS_START_DEFAULT}?)")
|
||||
feedback_history.append(f"Rejected: no out-of-sample data after {OOS_START_DEFAULT}.")
|
||||
progress.update(task, advance=1)
|
||||
continue
|
||||
|
||||
# Monte Carlo p-value (edge significance)
|
||||
mc_pvalue = bt_result.get('mc_pvalue')
|
||||
|
||||
# Rolling walk-forward metrics
|
||||
wf_consistency = bt_result.get('wf_oos_consistency')
|
||||
wf_sharpe_mean = bt_result.get('wf_oos_sharpe_mean')
|
||||
|
||||
# Check acceptance criteria — OOS must be profitable + statistically significant
|
||||
mc_ok = mc_pvalue is None or mc_pvalue < 0.20 # lenient: top 20% non-random
|
||||
wf_ok = wf_consistency is None or wf_consistency >= 0.5 # ≥50% of WF windows profitable
|
||||
if (abs(ic or 0) > MIN_IC and sharpe > MIN_SHARPE and trades > MIN_TRADES and dd > MAX_DRAWDOWN
|
||||
and oos_sharpe > 0.0 and oos_monthly > 0.0 and mc_ok and wf_ok):
|
||||
# ACCEPT
|
||||
strategy['real_backtest'] = bt_result
|
||||
strategy['metrics'] = bt_result
|
||||
@@ -414,6 +602,28 @@ def main(target_count=10):
|
||||
'annual_return_pct': bt_result.get('annual_return_pct', 0),
|
||||
'real_ic': ic, 'real_n_trades': trades, 'real_backtest_status': 'success',
|
||||
'n_bars': bt_result.get('n_bars', 0), 'n_months': bt_result.get('n_months', 0),
|
||||
'trading_style': TRADING_STYLE,
|
||||
'ohlcv_only': OHLCV_ONLY,
|
||||
'engine': 'ftmo_v2',
|
||||
'txn_cost_bps': TXN_COST_BPS,
|
||||
# Walk-forward OOS split
|
||||
'oos_sharpe': bt_result.get('oos_sharpe'),
|
||||
'oos_monthly_return_pct': bt_result.get('oos_monthly_return_pct'),
|
||||
'oos_max_drawdown': bt_result.get('oos_max_drawdown'),
|
||||
'oos_win_rate': bt_result.get('oos_win_rate'),
|
||||
'oos_n_trades': bt_result.get('oos_n_trades'),
|
||||
'is_sharpe': bt_result.get('is_sharpe'),
|
||||
'is_monthly_return_pct': bt_result.get('is_monthly_return_pct'),
|
||||
'oos_start': bt_result.get('oos_start'),
|
||||
# Rolling walk-forward
|
||||
'wf_n_windows': bt_result.get('wf_n_windows'),
|
||||
'wf_oos_sharpe_mean': wf_sharpe_mean,
|
||||
'wf_oos_sharpe_std': bt_result.get('wf_oos_sharpe_std'),
|
||||
'wf_oos_monthly_return_mean': bt_result.get('wf_oos_monthly_return_mean'),
|
||||
'wf_oos_consistency': wf_consistency,
|
||||
# Monte Carlo significance
|
||||
'mc_pvalue': mc_pvalue,
|
||||
'mc_n_permutations': bt_result.get('mc_n_permutations'),
|
||||
}
|
||||
|
||||
fname = f"{int(time.time())}_{strategy['strategy_name']}.json"
|
||||
@@ -435,8 +645,19 @@ def main(target_count=10):
|
||||
progress.console.print(f"[green]✓ Strategy #{len(accepted)}:[/green] {strategy['strategy_name']} "
|
||||
f"IC={ic:.4f}, Sharpe={sharpe:.3f}, Trades={trades}, DD={dd:.1%}")
|
||||
else:
|
||||
_log.info(f"REJECTED IC={ic:.4f} Sharpe={sharpe:.2f} Trades={trades} DD={dd:.1%}")
|
||||
feedback_history.append(f"Failed: IC={ic:.4f}, Sharpe={sharpe:.2f}, Trades={trades}, DD={dd:.1%}. Need |IC|>{MIN_IC}, Sharpe>{MIN_SHARPE}, Trades>{MIN_TRADES}")
|
||||
oos_info = f"OOS_Sharpe={oos_sharpe:+.2f} OOS_Mon={oos_monthly:+.2f}%" if oos_sharpe is not None else ""
|
||||
mc_info = f" MC_p={mc_pvalue:.2f}" if mc_pvalue is not None else ""
|
||||
wf_info = f" WF_consistency={wf_consistency:.0%}" if wf_consistency is not None else ""
|
||||
_ic = ic or 0; _sh = sharpe or 0; _dd = dd or 0
|
||||
_log.info(f"REJECTED IC={_ic:.4f} Sharpe={_sh:.2f} Trades={trades} DD={_dd:.1%} {oos_info}{mc_info}{wf_info}")
|
||||
feedback_history.append(
|
||||
f"Failed: IC={_ic:.4f}, Sharpe={_sh:.2f}, Trades={trades}, DD={_dd:.1%}, "
|
||||
f"OOS_Sharpe={oos_sharpe:+.2f}, OOS_Monthly={oos_monthly:+.2f}%"
|
||||
+ (f", MC_p={mc_pvalue:.2f}" if mc_pvalue is not None else "")
|
||||
+ (f", WF_consistency={wf_consistency:.0%}" if wf_consistency is not None else "")
|
||||
+ f". Need |IC|>{MIN_IC}, Sharpe>{MIN_SHARPE}, Trades>{MIN_TRADES}, "
|
||||
f"OOS_Sharpe>0, OOS_Monthly>0, MC_p<0.20, WF_consistency≥50%."
|
||||
)
|
||||
|
||||
progress.update(task, advance=1)
|
||||
|
||||
|
||||
@@ -22,9 +22,11 @@ from __future__ import annotations
|
||||
import argparse
|
||||
import csv
|
||||
import json
|
||||
import logging
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
@@ -34,13 +36,41 @@ from rich.console import Console
|
||||
from rich.progress import BarColumn, Progress, SpinnerColumn, TextColumn, TimeElapsedColumn
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
from rdagent.components.backtesting.vbt_backtest import backtest_signal # noqa: E402
|
||||
from rdagent.components.backtesting.vbt_backtest import backtest_signal_ftmo # noqa: E402
|
||||
|
||||
OHLCV_PATH = Path("/home/nico/Predix/git_ignore_folder/factor_implementation_source_data/intraday_pv.h5")
|
||||
FACTORS_VALUES_DIR = Path("/home/nico/Predix/results/factors/values")
|
||||
STRATEGIES_DIR = Path("/home/nico/Predix/results/strategies_new")
|
||||
|
||||
console = Console()
|
||||
# ── Logging setup: everything printed goes to log file + stdout ───────────────
|
||||
_LOG_DIR = Path(__file__).resolve().parent.parent / "git_ignore_folder" / "logs"
|
||||
_LOG_DIR.mkdir(parents=True, exist_ok=True)
|
||||
_log_file_path = _LOG_DIR / f"rebacktest_{datetime.now().strftime('%Y%m%d_%H%M%S')}.log"
|
||||
_log_file = open(_log_file_path, "w", encoding="utf-8", buffering=1)
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||||
handlers=[
|
||||
logging.StreamHandler(sys.stdout),
|
||||
logging.FileHandler(_log_file_path, encoding="utf-8"),
|
||||
],
|
||||
)
|
||||
|
||||
class _TeeFile:
|
||||
"""Writes to both stdout and log file — used as Rich Console file."""
|
||||
def __init__(self, *files):
|
||||
self._files = files
|
||||
def write(self, data):
|
||||
for f in self._files:
|
||||
f.write(data)
|
||||
def flush(self):
|
||||
for f in self._files:
|
||||
f.flush()
|
||||
def fileno(self):
|
||||
return self._files[0].fileno()
|
||||
|
||||
console = Console(file=_TeeFile(sys.stdout, _log_file), highlight=False)
|
||||
|
||||
|
||||
def load_close() -> pd.Series:
|
||||
@@ -154,11 +184,12 @@ def rebacktest_one(
|
||||
# Signal can arrive on either the factor index or the close index.
|
||||
signal = signal.reindex(close_a.index).ffill().fillna(0)
|
||||
|
||||
result = backtest_signal(
|
||||
result = backtest_signal_ftmo(
|
||||
close=close_a,
|
||||
signal=signal,
|
||||
txn_cost_bps=txn_cost_bps,
|
||||
freq="1min",
|
||||
wf_rolling=True,
|
||||
mc_n_permutations=200,
|
||||
)
|
||||
result["status_detail"] = result.pop("status")
|
||||
result["status"] = "ok"
|
||||
@@ -173,9 +204,13 @@ def main() -> None:
|
||||
help="Strategy directory to re-backtest")
|
||||
parser.add_argument("--csv", type=Path, default=None,
|
||||
help="Write a CSV report to this path")
|
||||
parser.add_argument("--txn-cost-bps", type=float, default=1.5)
|
||||
parser.add_argument("--txn-cost-bps", type=float, default=2.14,
|
||||
help="Transaction cost bps (default 2.14 ≈ 2.35 pip EUR/USD)")
|
||||
parser.add_argument("--write-back", action="store_true",
|
||||
help="Overwrite summary field in strategy JSON files with new results")
|
||||
args = parser.parse_args()
|
||||
|
||||
console.print(f"[dim]Log: {_log_file_path}[/dim]")
|
||||
console.print(f"[cyan]Loading OHLCV close...[/cyan]")
|
||||
close = load_close()
|
||||
console.print(f"[green]✓[/green] {len(close):,} 1-min bars "
|
||||
@@ -207,6 +242,51 @@ def main() -> None:
|
||||
|
||||
bt = rebacktest_one(data, close, args.txn_cost_bps)
|
||||
|
||||
if args.write_back and bt.get("status") == "ok":
|
||||
data["summary"] = {
|
||||
"sharpe": bt.get("sharpe"),
|
||||
"max_drawdown": bt.get("max_drawdown"),
|
||||
"win_rate": bt.get("win_rate"),
|
||||
"monthly_return_pct": bt.get("monthly_return_pct"),
|
||||
"real_ic": data.get("summary", {}).get("real_ic"),
|
||||
"real_n_trades": bt.get("n_trades"),
|
||||
"total_return": bt.get("total_return"),
|
||||
"annualized_return": bt.get("annualized_return"),
|
||||
"ftmo_daily_loss_hit": bt.get("ftmo_daily_loss_hit"),
|
||||
"ftmo_total_loss_hit": bt.get("ftmo_total_loss_hit"),
|
||||
"trading_style": data.get("summary", {}).get("trading_style"),
|
||||
"engine": "ftmo_v2",
|
||||
"txn_cost_bps": args.txn_cost_bps,
|
||||
# Walk-forward OOS
|
||||
"is_sharpe": bt.get("is_sharpe"),
|
||||
"is_monthly_return_pct": bt.get("is_monthly_return_pct"),
|
||||
"oos_sharpe": bt.get("oos_sharpe"),
|
||||
"oos_monthly_return_pct": bt.get("oos_monthly_return_pct"),
|
||||
"oos_max_drawdown": bt.get("oos_max_drawdown"),
|
||||
"oos_win_rate": bt.get("oos_win_rate"),
|
||||
"oos_n_trades": bt.get("oos_n_trades"),
|
||||
"oos_start": bt.get("oos_start"),
|
||||
# Rolling walk-forward
|
||||
"wf_n_windows": bt.get("wf_n_windows"),
|
||||
"wf_oos_sharpe_mean": bt.get("wf_oos_sharpe_mean"),
|
||||
"wf_oos_sharpe_std": bt.get("wf_oos_sharpe_std"),
|
||||
"wf_oos_monthly_return_mean": bt.get("wf_oos_monthly_return_mean"),
|
||||
"wf_oos_consistency": bt.get("wf_oos_consistency"),
|
||||
# Monte Carlo significance
|
||||
"mc_pvalue": bt.get("mc_pvalue"),
|
||||
"mc_n_permutations": bt.get("mc_n_permutations"),
|
||||
}
|
||||
data["sharpe_ratio"] = bt.get("sharpe")
|
||||
data["max_drawdown"] = bt.get("max_drawdown")
|
||||
data["win_rate"] = bt.get("win_rate")
|
||||
data["total_return"] = bt.get("total_return")
|
||||
data["reevaluation_status"] = "ftmo_v2"
|
||||
try:
|
||||
import json as _json
|
||||
f.write_text(_json.dumps(data, indent=2, ensure_ascii=False))
|
||||
except Exception as _e:
|
||||
logging.warning(f"write-back failed for {f.name}: {_e}")
|
||||
|
||||
row = {
|
||||
"file": f.name,
|
||||
"name": name,
|
||||
@@ -220,10 +300,23 @@ def main() -> None:
|
||||
"new_dd": bt.get("max_drawdown"),
|
||||
"new_trades": bt.get("n_trades"),
|
||||
"new_total_return": bt.get("total_return"),
|
||||
"new_monthly_pct": bt.get("monthly_return_pct"),
|
||||
"new_annual_return_cagr": None,
|
||||
"data_quality": bt.get("data_quality_flag"),
|
||||
# OOS walk-forward
|
||||
"is_sharpe": bt.get("is_sharpe"),
|
||||
"is_monthly_pct": bt.get("is_monthly_return_pct"),
|
||||
"oos_sharpe": bt.get("oos_sharpe"),
|
||||
"oos_monthly_pct": bt.get("oos_monthly_return_pct"),
|
||||
"oos_dd": bt.get("oos_max_drawdown"),
|
||||
"oos_trades": bt.get("oos_n_trades"),
|
||||
# Rolling walk-forward
|
||||
"wf_n_windows": bt.get("wf_n_windows"),
|
||||
"wf_oos_sharpe_mean": bt.get("wf_oos_sharpe_mean"),
|
||||
"wf_oos_consistency": bt.get("wf_oos_consistency"),
|
||||
# Monte Carlo
|
||||
"mc_pvalue": bt.get("mc_pvalue"),
|
||||
}
|
||||
# annualized CAGR is not in forward_returns wrapper; use annual_return_pct/100 proxy
|
||||
if "annualized_return" in bt:
|
||||
row["new_annual_return_cagr"] = bt["annualized_return"]
|
||||
rows.append(row)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user